"""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()), }