#!/usr/bin/env python3 """Run one source-free STRATA Native LM v1 authoritative read.""" from __future__ import annotations import argparse import json import os from pathlib import Path import torch from transformers import AutoModelForCausalLM, AutoTokenizer from strata.data.native_lm_integration import NativeLMExample, address_codes from strata.eval.native_lm_frame_separated_copy import frame_separated_generate from strata.memory_model.codec import FrozenMemoryCodec from strata.modeling.exact_payload_realizer import PayloadAuthority from strata.modeling.native_lm_integration import ( QualifiedP0M2Reader, StrataMemoryConditionedLM, ) from strata.modeling.structural_copy import StructuralCopyActionHead from strata.training.native_lm_integration import compact_state_table def load_model(root: Path, base_model: str, device: torch.device): config = json.loads((root / "configs/strata_native_lm_system_v1.json").read_text()) m1_config = json.loads( (root / "configs/strata_native_lm_integration_m1.json").read_text() ) model_config = m1_config["model"] tokenizer = AutoTokenizer.from_pretrained(base_model, local_files_only=True) if tokenizer.pad_token_id is None: tokenizer.pad_token = tokenizer.eos_token backbone = AutoModelForCausalLM.from_pretrained( base_model, local_files_only=True, torch_dtype=torch.bfloat16, attn_implementation=model_config["attention_implementation"], ).to(device) backbone.config.use_cache = False codec = FrozenMemoryCodec( checkpoint_path=root / "checkpoints/P0_M2_CHECKPOINT_FINAL.pt", config_path=root / "configs/strata_native_lm_p0_m2_v1.json", device="cpu", ) model = StrataMemoryConditionedLM( backbone, qualified_reader=QualifiedP0M2Reader(codec.model), layer_indices=model_config["memory_port_layers"], compact_width=int(model_config["compact_width"]), address_width=int(model_config["address_width"]), payload_width=int(m1_config["substrate"]["payload_width"]), memory_width=int(model_config["memory_width"]), memory_tokens=int(model_config["memory_tokens"]), attention_width=int(model_config["attention_width"]), heads=int(model_config["attention_heads"]), adapter_rank=int(model_config["adapter_rank"]), payload_classes=int(model_config["payload_classes"]), auxiliary_payload_loss_weight=float( model_config["auxiliary_payload_loss_weight"] ), ).to(device) checkpoint = torch.load( root / "checkpoints/MEMORY_PATH_FINAL.pt", map_location="cpu", weights_only=False, ) model.load_trainable_state_dict(checkpoint["state"]) model.eval() for parameter in model.parameters(): parameter.requires_grad_(False) head_checkpoint = torch.load( root / "checkpoints/ACTION_HEAD_FINAL.pt", map_location="cpu", weights_only=False, ) head = StructuralCopyActionHead(int(head_checkpoint["hidden_size"])).to(device) head.load_state_dict(head_checkpoint["state"], strict=True) head.eval() for parameter in head.parameters(): parameter.requires_grad_(False) return config, tokenizer, model, head, codec def main() -> None: parser = argparse.ArgumentParser() parser.add_argument( "--base-model", default=os.environ.get("STRATA_BASE_MODEL"), help="Local Qwen3-4B-Instruct-2507 snapshot", ) parser.add_argument("--event", required=True) parser.add_argument("--predicate", required=True) parser.add_argument("--role", required=True) parser.add_argument("--payload-handle", type=int, required=True) parser.add_argument("--payload", required=True) parser.add_argument("--query", required=True) parser.add_argument("--event-version", type=int, default=1) parser.add_argument("--device", default="cuda:0") args = parser.parse_args() if not args.base_model: parser.error("--base-model or STRATA_BASE_MODEL is required") if not 1 <= args.payload_handle <= 255: parser.error("--payload-handle must be in [1,255]") root = Path(__file__).resolve().parent device = torch.device(args.device) config, tokenizer, model, head, codec = load_model(root, args.base_model, device) row = NativeLMExample( example_id="release-request", split="release", schema=args.event.split(":", 1)[0], field=args.role, event=args.event, predicate=args.predicate, role=args.role, value_type="authoritative", payload_handle=args.payload_handle, value=args.payload, address_codes=address_codes(args.event, args.predicate, args.role), query=args.query, full_history_query=args.query, answer=f"The {args.role.replace('_', ' ')} is {args.payload}.", operation="point", age_windows=0, ) authority = PayloadAuthority.issue( event=args.event, predicate=args.predicate, role=args.role, handle=args.payload_handle, payload=args.payload, version=args.event_version, ) frame = config["frame"] outputs, timing = frame_separated_generate( model, head, tokenizer, [row], compact_state_table(codec), [[authority]], [0], batch_size=1, max_actions=int(config["evaluation"]["max_actions"]), frame_handle=int(frame["canonical_frame_handle"]), frame_surrogate=str(frame["canonical_frame_surrogate"]), terminator=str(frame["structural_terminator"]), current_versions=[args.event_version], ) result = outputs[0] print( json.dumps( { "answer": result.text, "frame": result.frame, "status": result.status, "payload_handle": result.controller_handle, "receipt": authority.receipt, "timing": timing, }, ensure_ascii=False, sort_keys=True, ) ) if __name__ == "__main__": main()