test_dmath.c 16 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436
  1. /* Per-function binary64 tests; see tests/dmath/README.md. No host-libm oracle. */
  2. #include "pocketpy/common/dmath.h"
  3. #include <fenv.h>
  4. #include <inttypes.h>
  5. #include <stdio.h>
  6. #include <stdlib.h>
  7. #include <string.h>
  8. typedef enum {
  9. T_isfinite,
  10. T_isinf,
  11. T_isnan,
  12. T_isnormal,
  13. T_fabs,
  14. T_copysign,
  15. T_fmin,
  16. T_fmax,
  17. T_ceil,
  18. T_floor,
  19. T_trunc,
  20. T_modf,
  21. T_fmod,
  22. T_sqrt,
  23. T_cbrt,
  24. T_exp,
  25. T_exp2,
  26. T_exp10,
  27. T_pow,
  28. T_log,
  29. T_log2,
  30. T_log10,
  31. T_log_base,
  32. T_sin,
  33. T_cos,
  34. T_tan,
  35. T_sincos,
  36. T_asin,
  37. T_acos,
  38. T_atan,
  39. T_atan2
  40. } Operation;
  41. typedef struct {
  42. const char* category;
  43. const char* name;
  44. Operation op;
  45. int outputs;
  46. int rounding_independent;
  47. } Function;
  48. const static Function functions[] = {
  49. {"classification", "isfinite", T_isfinite, 1, 1},
  50. {"classification", "isinf", T_isinf, 1, 1},
  51. {"classification", "isnan", T_isnan, 1, 1},
  52. {"classification", "isnormal", T_isnormal, 1, 1},
  53. {"sign_and_order", "fabs", T_fabs, 1, 1},
  54. {"sign_and_order", "copysign", T_copysign, 1, 1},
  55. {"sign_and_order", "fmin", T_fmin, 1, 1},
  56. {"sign_and_order", "fmax", T_fmax, 1, 1},
  57. {"rounding_and_remainder", "ceil", T_ceil, 1, 1},
  58. {"rounding_and_remainder", "floor", T_floor, 1, 1},
  59. {"rounding_and_remainder", "trunc", T_trunc, 1, 1},
  60. {"rounding_and_remainder", "modf", T_modf, 2, 1},
  61. {"rounding_and_remainder", "fmod", T_fmod, 1, 1},
  62. {"roots", "sqrt", T_sqrt, 1, 0},
  63. {"roots", "cbrt", T_cbrt, 1, 0},
  64. {"exponentials", "exp", T_exp, 1, 0},
  65. {"exponentials", "exp2", T_exp2, 1, 0},
  66. {"exponentials", "exp10", T_exp10, 1, 0},
  67. {"exponentials", "pow", T_pow, 1, 0},
  68. {"logarithms", "log", T_log, 1, 0},
  69. {"logarithms", "log2", T_log2, 1, 0},
  70. {"logarithms", "log10", T_log10, 1, 0},
  71. {"logarithms", "log_base", T_log_base, 1, 0},
  72. {"trigonometry", "sin", T_sin, 1, 0},
  73. {"trigonometry", "cos", T_cos, 1, 0},
  74. {"trigonometry", "tan", T_tan, 1, 0},
  75. {"trigonometry", "sincos", T_sincos, 2, 0},
  76. {"inverse_trigonometry", "asin", T_asin, 1, 0},
  77. {"inverse_trigonometry", "acos", T_acos, 1, 0},
  78. {"inverse_trigonometry", "atan", T_atan, 1, 0},
  79. {"inverse_trigonometry", "atan2", T_atan2, 1, 0},
  80. };
  81. const static uint64_t magnitude_mask = UINT64_C(0x7fffffffffffffff);
  82. const static uint64_t infinity_bits = UINT64_C(0x7ff0000000000000);
  83. const static uint64_t canonical_nan = UINT64_C(0x7ff8000000000000);
  84. const static uint64_t sign_bit = UINT64_C(0x8000000000000000);
  85. static int failures;
  86. static void evaluate(const Function* f, uint64_t a, uint64_t b, uint64_t out[2]) {
  87. double x = pk_dmath_from_bits(a), y = pk_dmath_from_bits(b), z = 0, aux = 0;
  88. out[1] = 0;
  89. switch(f->op) {
  90. case T_isfinite: out[0] = dmath_isfinite(x); return;
  91. case T_isinf: out[0] = dmath_isinf(x); return;
  92. case T_isnan: out[0] = dmath_isnan(x); return;
  93. case T_isnormal: out[0] = dmath_isnormal(x); return;
  94. case T_fabs: z = dmath_fabs(x); break;
  95. case T_copysign: z = dmath_copysign(x, y); break;
  96. case T_fmin: z = dmath_fmin(x, y); break;
  97. case T_fmax: z = dmath_fmax(x, y); break;
  98. case T_ceil: z = dmath_ceil(x); break;
  99. case T_floor: z = dmath_floor(x); break;
  100. case T_trunc: z = dmath_trunc(x); break;
  101. case T_modf: z = dmath_modf(x, &aux); break;
  102. case T_fmod: z = dmath_fmod(x, y); break;
  103. case T_sqrt: z = dmath_sqrt(x); break;
  104. case T_cbrt: z = dmath_cbrt(x); break;
  105. case T_exp: z = dmath_exp(x); break;
  106. case T_exp2: z = dmath_exp2(x); break;
  107. case T_exp10: z = dmath_exp10(x); break;
  108. case T_pow: z = dmath_pow(x, y); break;
  109. case T_log: z = dmath_log(x); break;
  110. case T_log2: z = dmath_log2(x); break;
  111. case T_log10: z = dmath_log10(x); break;
  112. case T_log_base: z = dmath_log_base(x, y); break;
  113. case T_sin: z = dmath_sin(x); break;
  114. case T_cos: z = dmath_cos(x); break;
  115. case T_tan: z = dmath_tan(x); break;
  116. case T_sincos: dmath_sincos(x, &z, &aux); break;
  117. case T_asin: z = dmath_asin(x); break;
  118. case T_acos: z = dmath_acos(x); break;
  119. case T_atan: z = dmath_atan(x); break;
  120. case T_atan2: z = dmath_atan2(x, y); break;
  121. }
  122. out[0] = pk_dmath_bits(z);
  123. out[1] = pk_dmath_bits(aux);
  124. }
  125. static void mismatch(const Function* f,
  126. const char* label,
  127. const char* check,
  128. uint64_t a,
  129. uint64_t b,
  130. uint64_t actual,
  131. uint64_t expected) {
  132. if(failures++ < 30)
  133. fprintf(stderr,
  134. "FAIL dmath_%s/%s [%s] x=%016" PRIx64 " y=%016" PRIx64 " got=%016" PRIx64
  135. " expected=%016" PRIx64 "\n",
  136. f->name,
  137. label,
  138. check,
  139. a,
  140. b,
  141. actual,
  142. expected);
  143. }
  144. static int accurate(uint64_t actual, uint64_t reference, unsigned ulps) {
  145. if(actual == reference) return 1;
  146. uint64_t a = actual & magnitude_mask, r = reference & magnitude_mask;
  147. // Nonfinite outputs have semantic checks, never NaN payload/sign checks.
  148. if(r > infinity_bits) return a > infinity_bits;
  149. if(r == infinity_bits) return a == infinity_bits && !((actual ^ reference) & sign_bit);
  150. if(r == 0 || a >= infinity_bits || ((actual ^ reference) & sign_bit)) return 0;
  151. return (a > r ? a - r : r - a) <= ulps;
  152. }
  153. static uint64_t signature_word(uint64_t u) {
  154. // A fingerprint records NaN classification, not its representation.
  155. return (u & magnitude_mask) > infinity_bits ? canonical_nan : u;
  156. }
  157. static uint64_t parse_word(const char* text) {
  158. char* end;
  159. uint64_t u = strtoull(text, &end, 16);
  160. if(strlen(text) != 16 || *end) {
  161. fprintf(stderr, "Invalid binary64 word: %s\n", text);
  162. exit(1);
  163. }
  164. return u;
  165. }
  166. static uint64_t parse_result(const char* text) {
  167. if(strcmp(text, "nan") == 0) return canonical_nan;
  168. if(strcmp(text, "+inf") == 0) return infinity_bits;
  169. if(strcmp(text, "-inf") == 0) return infinity_bits | sign_bit;
  170. return parse_word(text);
  171. }
  172. static unsigned
  173. named_cases(const Function* f, const char* directory, int probe, uint64_t* expected_sweep) {
  174. char path[1024], line[1024];
  175. snprintf(path, sizeof(path), "%s/%s.txt", directory, f->name);
  176. FILE* file = fopen(path, "r");
  177. if(!file) {
  178. perror(path);
  179. exit(1);
  180. }
  181. unsigned count = 0;
  182. int found_sweep = 0;
  183. while(fgets(line, sizeof(line), file)) {
  184. char fingerprint[32];
  185. if(sscanf(line, "# sweep %31s", fingerprint) == 1) {
  186. found_sweep = 1;
  187. if(strcmp(fingerprint, "PENDING") == 0) {
  188. if(!probe) {
  189. fprintf(stderr, "%s: unreviewed sweep\n", path);
  190. failures++;
  191. }
  192. } else
  193. *expected_sweep = parse_word(fingerprint);
  194. continue;
  195. }
  196. if(line[0] == '#' || line[0] == '\n' || line[0] == '\r') continue;
  197. char label[128], words[6][32];
  198. unsigned ulps;
  199. if(sscanf(line,
  200. "%127s %31s %31s %31s %31s %31s %31s %u",
  201. label,
  202. words[0],
  203. words[1],
  204. words[2],
  205. words[3],
  206. words[4],
  207. words[5],
  208. &ulps) != 8) {
  209. fprintf(stderr, "Malformed case in %s\n", path);
  210. exit(1);
  211. }
  212. uint64_t a = parse_word(words[0]), b = parse_word(words[1]), out[2];
  213. evaluate(f, a, b, out);
  214. for(int i = 0; i < f->outputs; i++) {
  215. uint64_t ref = parse_result(words[4 + i]);
  216. if(!accurate(out[i], ref, ulps))
  217. mismatch(f, label, i ? "accuracy/output1" : "accuracy/output0", a, b, out[i], ref);
  218. if(!probe) {
  219. if(strcmp(words[2 + i], "PENDING") == 0) {
  220. fprintf(stderr, "%s/%s: unreviewed result\n", f->name, label);
  221. failures++;
  222. } else {
  223. uint64_t want = parse_result(words[2 + i]);
  224. if(!accurate(out[i], want, 0))
  225. mismatch(f, label, i ? "bits/output1" : "bits/output0", a, b, out[i], want);
  226. }
  227. }
  228. }
  229. if(probe == 1)
  230. printf("CASE %s %s %016" PRIx64 " %016" PRIx64 "\n",
  231. f->name,
  232. label,
  233. signature_word(out[0]),
  234. signature_word(out[1]));
  235. count++;
  236. }
  237. if(ferror(file) || !count || !found_sweep) {
  238. fprintf(stderr, "Incomplete corpus: %s\n", path);
  239. exit(1);
  240. }
  241. fclose(file);
  242. return count;
  243. }
  244. static uint64_t hash_word(uint64_t hash, uint64_t word) {
  245. for(int byte = 0; byte < 8; byte++, word >>= 8)
  246. hash = (hash ^ (word & 255)) * UINT64_C(1099511628211);
  247. return hash;
  248. }
  249. static uint64_t random_word(uint64_t* state) {
  250. // SplitMix64, with unsigned wrapping arithmetic and a new seed per function.
  251. uint64_t z = (*state += UINT64_C(0x9e3779b97f4a7c15));
  252. z = (z ^ (z >> 30)) * UINT64_C(0xbf58476d1ce4e5b9);
  253. z = (z ^ (z >> 27)) * UINT64_C(0x94d049bb133111eb);
  254. return z ^ (z >> 31);
  255. }
  256. static void invariants(const Function* f, uint64_t a, uint64_t b, const uint64_t out[2]) {
  257. uint64_t ax = a & magnitude_mask, ay = b & magnitude_mask, expected;
  258. int has_exact = 1;
  259. switch(f->op) {
  260. case T_isfinite: expected = ax < infinity_bits; break;
  261. case T_isinf: expected = ax == infinity_bits; break;
  262. case T_isnan: expected = ax > infinity_bits; break;
  263. case T_isnormal: expected = ax >= UINT64_C(0x0010000000000000) && ax < infinity_bits; break;
  264. case T_fabs: expected = ax; break;
  265. case T_copysign: expected = ax | (b & sign_bit); break;
  266. default:
  267. has_exact = 0;
  268. expected = 0;
  269. break;
  270. }
  271. if(has_exact) {
  272. if(!accurate(out[0], expected, 0))
  273. mismatch(f, "sweep", "finite bits/nonfinite class", a, b, out[0], expected);
  274. return;
  275. }
  276. if(f->op == T_fmin || f->op == T_fmax) {
  277. uint64_t reverse[2];
  278. evaluate(f, b, a, reverse);
  279. if(!accurate(out[0], reverse[0], 0))
  280. mismatch(f, "sweep", "commutativity", a, b, out[0], reverse[0]);
  281. }
  282. if(f->op == T_sincos) {
  283. expected = pk_dmath_bits(dmath_sin(pk_dmath_from_bits(a)));
  284. if(!accurate(out[0], expected, 0))
  285. mismatch(f, "sweep", "separate sine", a, b, out[0], expected);
  286. expected = pk_dmath_bits(dmath_cos(pk_dmath_from_bits(a)));
  287. if(!accurate(out[1], expected, 0))
  288. mismatch(f, "sweep", "separate cosine", a, b, out[1], expected);
  289. }
  290. if(f->op == T_fmod && ax < infinity_bits && ay != 0 && ay <= infinity_bits) {
  291. if((out[0] & magnitude_mask) >= ay || ((a ^ out[0]) & sign_bit))
  292. mismatch(f, "sweep", "remainder range/sign", a, b, out[0], a & sign_bit);
  293. }
  294. if(f->op == T_modf && ax < infinity_bits) {
  295. double fraction = pk_dmath_from_bits(out[0]), integral = pk_dmath_from_bits(out[1]);
  296. expected = pk_dmath_bits(fraction + integral);
  297. if(expected != a) mismatch(f, "sweep", "recomposition", a, b, expected, a);
  298. if((out[0] & magnitude_mask) >= UINT64_C(0x3ff0000000000000) || ((out[0] ^ a) & sign_bit) ||
  299. ((out[1] ^ a) & sign_bit))
  300. mismatch(f, "sweep", "fraction range/sign", a, b, out[0], a & sign_bit);
  301. }
  302. }
  303. static uint64_t sweep(const Function* f) {
  304. uint64_t state = UINT64_C(0x243f6a8885a308d3);
  305. for(const char* p = f->name; *p; p++)
  306. state = hash_word(state, (unsigned char)*p);
  307. uint64_t hash = UINT64_C(14695981039346656037);
  308. // All exponent fields, with fresh significands of both signs, followed by
  309. // 16,384 full-range words and an additional stream in useful finite domains.
  310. for(unsigned i = 0; i < 4096 + 16384; i++) {
  311. uint64_t a = random_word(&state), b = random_word(&state), out[2];
  312. if(i < 4096)
  313. a = ((uint64_t)(i & 1) << 63) | ((uint64_t)(i / 2) << 52) |
  314. (a & UINT64_C(0x000fffffffffffff));
  315. evaluate(f, a, b, out);
  316. invariants(f, a, b, out);
  317. for(int j = 0; j < f->outputs; j++)
  318. hash = hash_word(hash, signature_word(out[j]));
  319. if(f->op >= T_exp && f->op <= T_pow) {
  320. double bounded = (double)(a >> 32) * 0x1p-21 - 1024.0;
  321. if(f->op == T_exp10) bounded *= 0.25;
  322. if(f->op == T_pow) {
  323. a = UINT64_C(0x3ff0000000000000) + (a & 4095);
  324. b = pk_dmath_bits(bounded * 0x1p40);
  325. } else
  326. a = pk_dmath_bits(bounded);
  327. } else if(f->op == T_asin || f->op == T_acos) {
  328. a = pk_dmath_bits((double)(a >> 11) * 0x1p-52 - 1.0);
  329. } else if(f->op >= T_log && f->op <= T_log_base) {
  330. a &= magnitude_mask;
  331. b &= magnitude_mask;
  332. } else
  333. continue;
  334. evaluate(f, a, b, out);
  335. invariants(f, a, b, out);
  336. for(int j = 0; j < f->outputs; j++)
  337. hash = hash_word(hash, signature_word(out[j]));
  338. }
  339. return hash;
  340. }
  341. static void run(const Function* f, const char* directory, int probe) {
  342. int before = failures;
  343. uint64_t frozen = 0;
  344. unsigned count = named_cases(f, directory, probe, &frozen);
  345. uint64_t actual = sweep(f);
  346. if(probe)
  347. printf("SWEEP %s %016" PRIx64 "\n", f->name, actual);
  348. else if(actual != frozen)
  349. mismatch(f, "sweep", "fingerprint", 0, 0, actual, frozen);
  350. if(f->rounding_independent) {
  351. const int modes[] = {FE_DOWNWARD, FE_UPWARD, FE_TOWARDZERO};
  352. for(unsigned i = 0; i < sizeof(modes) / sizeof(*modes); i++) {
  353. if(fesetround(modes[i]) != 0) {
  354. fputs("fesetround failed\n", stderr);
  355. exit(1);
  356. }
  357. named_cases(f, directory, probe ? 2 : 0, &frozen);
  358. uint64_t directed = sweep(f);
  359. if(directed != actual)
  360. mismatch(f, "sweep", "rounding mode dependence", modes[i], 0, directed, actual);
  361. }
  362. if(fesetround(FE_TONEAREST) != 0) {
  363. fputs("cannot restore rounding\n", stderr);
  364. exit(1);
  365. }
  366. }
  367. printf("%s %s/%s: %u named cases, sweep=%016" PRIx64 ", %d rounding mode(s)\n",
  368. failures == before ? "PASS" : "FAIL",
  369. f->category,
  370. f->name,
  371. count,
  372. actual,
  373. f->rounding_independent ? 4 : 1);
  374. }
  375. int main(int argc, char** argv) {
  376. const char* directory = "tests/dmath/cases";
  377. const char* only = NULL;
  378. int probe = 0, list = 0, ran = 0;
  379. for(int i = 1; i < argc; i++) {
  380. if(strcmp(argv[i], "--cases") == 0 && i + 1 < argc)
  381. directory = argv[++i];
  382. else if(strcmp(argv[i], "--probe") == 0)
  383. probe = 1;
  384. else if(strcmp(argv[i], "--list") == 0)
  385. list = 1;
  386. else if(argv[i][0] != '-' && !only)
  387. only = argv[i];
  388. else {
  389. fputs("Usage: test_dmath [--cases DIR] [--list] [--probe] [FUNCTION]\n", stderr);
  390. return 1;
  391. }
  392. }
  393. if(fegetround() != FE_TONEAREST) {
  394. fputs("dmath requires nearest-even\n", stderr);
  395. return 1;
  396. }
  397. volatile double tiny = pk_dmath_from_bits(3);
  398. volatile double normal = pk_dmath_from_bits(UINT64_C(0x0010000000000000));
  399. if(pk_dmath_bits(tiny + tiny) != 6 ||
  400. pk_dmath_bits(normal * 0.25) != UINT64_C(0x0004000000000000)) {
  401. fputs("dmath requires gradual underflow (FTZ/DAZ disabled)\n", stderr);
  402. return 1;
  403. }
  404. for(unsigned i = 0; i < sizeof(functions) / sizeof(*functions); i++) {
  405. const Function* f = &functions[i];
  406. if(only && strcmp(f->name, only) != 0) continue;
  407. if(list)
  408. printf("%s/%s\n", f->category, f->name);
  409. else
  410. run(f, directory, probe);
  411. ran++;
  412. }
  413. if(!ran) {
  414. fprintf(stderr, "Unknown dmath function: %s\n", only);
  415. return 1;
  416. }
  417. if(failures) fprintf(stderr, "%d dmath check(s) failed\n", failures);
  418. return failures ? 1 : 0;
  419. }