Xunzhuo's picture
Release Decision 1.0 Lex: typed decision specialist
92b5c9a verified
Raw History Blame
2.37 kB
"""Complete-input 1K inference using the bundled, unchanged native runtime."""
import argparse
import json
from pathlib import Path
NATIVE_MANIFEST_SHA256 = 'f288d873999832a3f37c6a7c4268c2ab309691e621794dbf7acab891acbbb7e6'
def main():
p = argparse.ArgumentParser(description='Decision Lex: complete1024, FP32 AMD inference')
p.add_argument('--native', default=str(Path(__file__).resolve().parent / 'native'))
p.add_argument('--manifest-sha256', default=NATIVE_MANIFEST_SHA256)
p.add_argument('--input', required=True, help='JSONL with id, state_text, question')
p.add_argument('--output', required=True, help='Fresh prediction JSONL; existing files are not overwritten')
p.add_argument('--batch-size', type=int, choices=range(1, 9), default=8)
args = p.parse_args()
records = []
with Path(args.input).open(encoding='utf-8') as stream:
for line in stream:
row = json.loads(line)
if set(row) != {'id', 'state_text', 'question'}:
raise ValueError('Inference rows must contain exactly id, state_text and question')
records.append(row)
if not records or len({r['id'] for r in records}) != len(records):
raise ValueError('Nonempty input with unique record IDs required')
if Path(args.output).exists():
raise ValueError('Prediction output must be fresh')
import torch
from decision_runtime import load_native
from decision_inference import predict_1k
if torch.version.hip is None or not torch.cuda.is_available() or torch.cuda.device_count() != 1:
raise RuntimeError('Expose exactly one AMD ROCm GPU; this package has no CPU/NVIDIA fallback')
torch.cuda.set_device(0); torch.set_num_threads(2)
torch.backends.cuda.matmul.allow_tf32 = False
torch.backends.cudnn.allow_tf32 = False
torch.backends.mha.set_fastpath_enabled(False)
native = load_native(args.native, expected_manifest_sha256=args.manifest_sha256, device='cuda:0')
answers = predict_1k(native, records, batch_size=args.batch_size)
if len(answers) != len(records):
raise RuntimeError('Incomplete inference result')
with Path(args.output).open('x', encoding='utf-8') as stream:
for answer in answers:
stream.write(json.dumps(answer, ensure_ascii=False, allow_nan=False) + '\n')
if __name__ == '__main__': main()