NanoJev-Web / python /predictor.py
candypunk's picture
Release NanoJev-Web browser action model
a0a9254 verified
Raw History Blame Contribute Delete
10.8 kB
"""Local browser decision inference. No downloads, teacher or game runtime."""
from contextlib import nullcontext
import json,math,os
from pathlib import Path
from schema import validate_request
def unique_object(pairs):
obj = {}
for key, value in pairs:
if key in obj:
raise ValueError(f"JSON contains duplicate key: {key}")
obj[key] = value
return obj
def reject_nonfinite(value):
raise ValueError(f"JSON disallows non-finite number: {value}")
def read_json(path):
return json.loads(Path(path).read_text(encoding="utf-8"), object_pairs_hook=unique_object,
parse_constant=reject_nonfinite)
def nonempty_text(value):
return isinstance(value, str) and bool(value.strip())
def prepare_examples(payload, tokenizer, max_length):
"""Segment encoding, candidate text and EOS follow the upstream training tokenization exactly."""
states = validate_request(payload)
if type(max_length) is not int or max_length <= 0:
raise ValueError("max_length must be a positive integer")
if type(tokenizer.eos_token_id) is not int or tokenizer.eos_token_id < 0:
raise ValueError("Checkpoint tokenizer must have a valid eos_token_id")
examples = []
for row in states:
for qid, q in row["questions"].items():
typ = q["type"]
ids = list(q["criteria"])
texts = [f"{key}: {q['criteria'][key]}" for key in ids]
segments = [f"State:\n{row['state']}\n",
f"Question type: {typ}\nQuestion:\n{q['instructions']}\n"]
prefix = sum([tokenizer.encode(t, add_special_tokens=False) for t in segments], [])
leaves = [prefix + tokenizer.encode(f"Candidate:\n{t}\nDecision:", add_special_tokens=False)
+ [tokenizer.eos_token_id] for t in texts]
largest = max(map(len, leaves))
if largest > max_length:
raise ValueError(f"{row['id']}:{qid} candidate path has {largest} tokens, exceeding max_length={max_length}; input was not truncated")
examples.append({"id": f"{row['id']}:{qid}", "state_id": row["id"], "qid": qid,
"type": typ, "candidate_ids": ids, "candidate_texts": texts,
"leaf_tokens": leaves})
return examples
def complete_question_batches(examples, batch_questions=0):
if type(batch_questions) is not int or batch_questions < 0:
raise ValueError("batch_questions must be nonnegative; zero batches all questions together")
size = batch_questions or len(examples)
if not examples:
return []
return [examples[i:i + size] for i in range(0, len(examples), size)]
def answer_from_probabilities(example, probabilities):
ids = example["candidate_ids"]
if len(probabilities) != len(ids) or not all(math.isfinite(p) and 0 <= p <= 1 for p in probabilities):
raise ValueError("Model produced invalid probabilities")
if abs(math.fsum(probabilities) - 1.0) > 1e-5:
raise ValueError("Model probabilities do not sum to one")
best = max(range(len(ids)), key=probabilities.__getitem__)
result = {"type": example["type"], "probabilities": dict(zip(ids, probabilities))}
result.update(choice=ids[best], value=ids[best])
return result
def local_checkpoint_files(checkpoint_dir):
root = Path(checkpoint_dir).expanduser().resolve(strict=True)
if not root.is_dir():
raise ValueError("Checkpoint must be a local directory")
paths = {"run_config": root / "config.json", "body_config": root / "backbone_config",
"tokenizer": root / "tokenizer", "weights": root / "best.safetensors"}
for label, path in paths.items():
if not path.exists():
raise ValueError(f"Checkpoint missing {label}: {path.name}")
if not paths["run_config"].is_file() or not paths["weights"].is_file():
raise ValueError("config.json and best.safetensors must be files")
if not paths["body_config"].is_dir() or not paths["tokenizer"].is_dir():
raise ValueError("backbone_config and tokenizer must be directories")
return root, paths
class DecisionPredictor:
"""Persistent local inference: load weights once, evaluate complete candidate sets per question."""
def __init__(self, checkpoint_dir, max_length=None, device_name="mps",
disable_native_triton=False, precision="fp32"):
if precision != "fp32":
raise ValueError("precision must be fp32")
root, paths = local_checkpoint_files(checkpoint_dir)
run_config = read_json(paths["run_config"])
if not isinstance(run_config, dict) or run_config.get("set_head") not in {"none", "attention"}:
raise ValueError("Checkpoint config has no valid set_head")
# Load local checkpoint only; no environment files, Hub requests or model downloads.
os.environ["HF_HUB_OFFLINE"] = "1"
os.environ["TRANSFORMERS_OFFLINE"] = "1"
os.environ["HF_HUB_DISABLE_TELEMETRY"] = "1"
import torch
from safetensors.torch import load_file
from transformers import AutoConfig, AutoModel, AutoTokenizer
device = torch.device(device_name)
if device.type == "mps":
if not torch.backends.mps.is_available():
raise ValueError("MPS is unavailable; native Apple Silicon Python is required")
elif device.type != "cpu":
raise ValueError("This package supports mps and cpu")
if precision != "fp32":
raise ValueError("This package uses FP32 inference")
tokenizer = AutoTokenizer.from_pretrained(str(paths["tokenizer"]), local_files_only=True,
trust_remote_code=False)
if tokenizer.pad_token_id is None:
tokenizer.pad_token = tokenizer.eos_token
body_config = AutoConfig.from_pretrained(str(paths["body_config"]), local_files_only=True,
trust_remote_code=False)
body_config.use_cache = False
limit = run_config.get("max_length", 512) if max_length is None else max_length
if type(limit) is not int or limit <= 0:
raise ValueError("max-length must be a positive integer")
context_limit = getattr(body_config, "max_position_embeddings", None)
if isinstance(context_limit, int) and limit > context_limit:
raise ValueError("max-length exceeds the backbone context length")
# Build the architecture from config; load every parameter from best.safetensors without downloading a base model.
body = AutoModel.from_config(body_config, attn_implementation="sdpa", trust_remote_code=False).float()
from model import DecisionModel
model = DecisionModel(body, run_config["set_head"])
weights = load_file(str(paths["weights"]), device="cpu")
model.load_state_dict(weights, strict=True)
del weights
model.to(device=device, dtype=torch.float32)
model.eval()
self.model = model
self.tokenizer = tokenizer
self.root = root
self.run_config = run_config
self.limit = limit
self.device = device
self.precision = precision
self.disable_native_triton = disable_native_triton
self.inference_calls = 0
self._torch = torch
def predict(self, payload, batch_questions=0, temperature=1.0):
states = validate_request(payload)
if not isinstance(temperature, (int, float)) or isinstance(temperature, bool) or not math.isfinite(temperature) or temperature <= 0:
raise ValueError("temperature must be finite and positive")
torch = self._torch
model, tokenizer = self.model, self.tokenizer
root, run_config, limit = self.root, self.run_config, self.limit
device, precision = self.device, self.precision
disable_native_triton = self.disable_native_triton
examples = prepare_examples(payload, tokenizer, limit)
batches = complete_question_batches(examples, batch_questions)
self.inference_calls += 1
model.eval()
outputs = {state["id"]: {"id": state["id"], "answers": {}} for state in states}
with torch.inference_mode():
for batch in batches:
context = nullcontext()
with context:
logits, _ = model(batch, tokenizer.pad_token_id)
for example, values in zip(batch, logits):
k = len(example["candidate_ids"])
scores = values[:k].float()
if not torch.isfinite(scores).all():
raise ValueError("Model produced non-finite logits; no partial prediction returned")
probabilities = (scores / temperature).softmax(-1).cpu().tolist()
outputs[example["state_id"]]["answers"][example["qid"]] = answer_from_probabilities(example, probabilities)
return {
"schema_version": "nanojev-browser-inference-v1",
"checkpoint": {"id": "NanoJev-Web/browser-head-v5", "base_model": run_config.get("model"),
"base_revision": run_config.get("resolved_model_revision"), "set_head": run_config["set_head"]},
"temperature": {"value": float(temperature), "fitted_by_this_command": False,
"note": "Applies the supplied scalar; the default of 1 does not imply calibration."},
"execution": {"device": str(device), "parameter_storage": "float32", "precision": precision,
"forward_autocast": "bfloat16" if precision == "bf16" else "disabled",
"states": len(states), "questions": len(examples),
"candidate_paths": sum(len(ex["leaf_tokens"]) for ex in examples),
"input_tokens": sum(len(path) for ex in examples for path in ex["leaf_tokens"]),
"padded_input_tokens": sum(
max(len(path) for ex in batch for path in ex["leaf_tokens"])
* sum(len(ex["leaf_tokens"]) for ex in batch) for batch in batches),
"output_tokens": 0,
"forward_passes": len(batches), "batch_questions_limit": batch_questions or "all",
"autoregressive_decode_steps": 0, "prefix_sharing": False,
"max_length": limit, "disable_native_triton": disable_native_triton,
"network_model_calls": 0, "persistent_model_load_count": 1,
"inference_call_index": self.inference_calls},
"states": list(outputs.values()),
}