generate.py 4.4 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104
  1. """Generate/check per-function cases from the independent oracle.
  2. Run: python scripts/dmath/generate.py [--check]
  3. Existing frozen output bits are preserved only if the case name and inputs
  4. match. New cases are emitted as PENDING and fail normal test runs until their
  5. results have been reviewed on independent compiler builds. This script never
  6. calls the implementation or silently blesses changed results.
  7. """
  8. import argparse
  9. from decimal import localcontext, InvalidOperation, DivisionByZero, Overflow
  10. from pathlib import Path
  11. import re
  12. import cases
  13. from oracle import reference, value, INF, MASK, SIGN
  14. ROOT = Path(__file__).resolve().parents[2]
  15. DEST = ROOT / 'tests/dmath/cases'
  16. def result_text(u):
  17. magnitude = u & MASK
  18. if magnitude > INF:
  19. return 'nan'
  20. if magnitude == INF:
  21. return '-inf' if u & SIGN else '+inf'
  22. return '%016x' % u
  23. def old_results(path):
  24. rows, sweep = {}, 'PENDING'
  25. if path.exists():
  26. for line in path.read_text().splitlines():
  27. if line.startswith('# sweep '):
  28. sweep = line.split()[2]
  29. elif line and not line.startswith('#'):
  30. fields = line.split()
  31. rows[tuple(fields[:3])] = [result_text(int(x, 16)) if len(x) == 16 else x
  32. for x in fields[3:5]]
  33. return rows, sweep
  34. def render(name, rows, path):
  35. existing, sweep = old_results(path)
  36. text = ['# dmath_' + name,
  37. '# Independent oracle: Decimal at 430 and 570 digits / exact Fraction arithmetic.',
  38. '# Frozen bits and sweep require review; generation never recalibrates them.',
  39. '# Nonfinite outputs: nan checks classification; +/-inf checks classification and sign.',
  40. '# sweep ' + sweep,
  41. '# case input0 input1 frozen0 frozen1 reference0 reference1 max_ulp']
  42. for row in rows:
  43. refs = []
  44. for precision in (430, 570):
  45. with localcontext() as context:
  46. context.prec = precision
  47. for trap in (InvalidOperation, DivisionByZero, Overflow):
  48. context.traps[trap] = False
  49. refs.append(reference(name, row.x, row.y))
  50. if refs[0] != refs[1]:
  51. raise RuntimeError('reference did not converge: ' + name + '/' + row.name)
  52. outputs = refs[0] + (0,) * (2 - len(refs[0]))
  53. key = (row.name, '%016x' % row.x, '%016x' % row.y)
  54. frozen = existing.get(key, ['PENDING', 'PENDING'])
  55. # Human-readable inputs stay next to their exact binary representations.
  56. comment = ' # x=' + value(row.x).hex() + ' y=' + value(row.y).hex()
  57. fields = [*key, *frozen, *[result_text(r) for r in outputs], str(cases.ULPS[name])]
  58. text.append(' '.join(fields) + comment)
  59. return '\n'.join(text) + '\n'
  60. def main():
  61. parser = argparse.ArgumentParser(description=__doc__)
  62. parser.add_argument('--check', action='store_true')
  63. args = parser.parse_args()
  64. exported = set(re.findall(r'^(?:double|int|void) dmath_(\w+)\(',
  65. (ROOT / 'include/pocketpy/common/dmath.h').read_text(), re.M))
  66. designed = {n for names in cases.GROUPS.values() for n in names}
  67. if exported != designed:
  68. raise RuntimeError('missing/extra function groups: ' + repr(exported ^ designed))
  69. registered = set(re.findall(r'\{\s*"[a-z_]+",\s*"(\w+)",\s*T_',
  70. (ROOT / 'src2/test_dmath.c').read_text()))
  71. if registered != designed:
  72. raise RuntimeError('C runner groups differ: ' + repr(registered ^ designed))
  73. DEST.mkdir(parents=True, exist_ok=True)
  74. total = 0
  75. for group, names in cases.GROUPS.items():
  76. for name in names:
  77. rows = getattr(cases, 'cases_' + name)()
  78. if len({row.name for row in rows}) != len(rows):
  79. raise RuntimeError('duplicate case name in ' + name)
  80. path = DEST / (name + '.txt')
  81. content = render(name, rows, path)
  82. if args.check:
  83. if not path.exists() or path.read_text() != content:
  84. raise RuntimeError(str(path) + ' is out of date')
  85. else:
  86. path.write_text(content, encoding='ascii', newline='\n')
  87. total += len(rows)
  88. print(group + '/' + name + ': ' + str(len(rows)) + ' cases', flush=True)
  89. print(str(len(designed)) + ' functions, ' + str(total) + ' named cases')
  90. if __name__ == '__main__':
  91. main()