Xunzhuo's picture
Release measured v1.2 update
d330a0b verified
Raw History Blame
3.58 kB
"""Run the real model-card request and save outputs; no invented predictions."""
import argparse
import hashlib
import json
from pathlib import Path
from .model import DecisionModel
REQUEST = {
'state': 'The customer reports that the same invoice was charged twice. They ask for a refund. There is no product outage.',
'questions': {
'destination': {'type': 'choice', 'instructions': 'Choose the team that handles this request.',
'criteria': {'billing': 'Invoices, payments, refunds and duplicate charges',
'technical': 'Product errors and troubleshooting'}},
'refund_requested': {'type': 'noul', 'instructions': 'Does the customer explicitly ask for a refund?'},
'urgency': {'type': 'score', 'instructions': 'Rate urgency using only these ordered levels.',
'criteria': ['Routine information request with no payment problem or outage',
'A payment or billing problem, with no product outage',
'An active product outage stopping the customer from working']},
},
}
def main():
parser = argparse.ArgumentParser()
parser.add_argument('model', help='Exported local directory or Hugging Face repo ID')
parser.add_argument('--revision', help='Required for a repo ID; prefer a full commit SHA')
parser.add_argument('--device', default='cuda:0')
parser.add_argument('--local-files-only', action='store_true')
parser.add_argument('--allow-unvalidated-runtime', action='store_true')
parser.add_argument('--output', type=Path, required=True)
args = parser.parse_args()
model = DecisionModel.from_pretrained(args.model, revision=args.revision, device=args.device,
local_files_only=args.local_files_only,
allow_unvalidated_runtime=args.allow_unvalidated_runtime)
response = model.decide(**REQUEST)
# Prove wrapper pass-through against the unchanged engine on this exact request.
reference = model._engine.decide(**REQUEST)
if response != reference:
raise AssertionError('Wrapper/direct-engine response mismatch')
overflow_message = None
try:
model.decide('overflow-test ' * 20000, {'check': {'type': 'noul', 'instructions': 'Is this a test?'}})
except ValueError as exc:
if 'no truncation allowed' not in str(exc):
raise
overflow_message = str(exc)
if overflow_message is None:
raise AssertionError('Oversized complete input was not rejected')
record = {'request': REQUEST, 'response': response, 'direct_engine_exact_response': True,
'overflow_rejected': True, 'overflow_message': overflow_message,
'bundle_manifest_sha256': hashlib.sha256((model.bundle_path / 'bundle-manifest.json').read_bytes()).hexdigest(),
'runtime': model.runtime, 'revision': args.revision, 'device': args.device,
'model_name': response['model'], 'example_source_sha256': hashlib.sha256(Path(__file__).read_bytes()).hexdigest()}
if model.runtime.get('normalization_profile') is not None:
import sys
record['normalization_telemetry'] = dict(sys.modules['_decision_process_normalization_profile_v1'].telemetry)
args.output.parent.mkdir(parents=True, exist_ok=True)
args.output.write_text(json.dumps(record, ensure_ascii=False, indent=2) + '\n')
print(json.dumps(response, ensure_ascii=False, indent=2))
if __name__ == '__main__':
main()