strata-native-lm / load_and_answer.py
nur-dev's picture
STRATA Native LM: non-commercial release
1d60d59 verified
Raw History Blame Contribute Delete
6.21 kB
#!/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()