Download python/predictor.py from candypunk/NanoJev-Web: direct link, hf CLI and curl.
- Browser
- Download file 10.8 kB
-
https://huggingface.co/candypunk/NanoJev-Web/resolve/main/python/predictor.py
- Command line
-
hf download hf://candypunk/NanoJev-Web/python/predictor.py
-
curl -L -o predictor.py https://huggingface.co/candypunk/NanoJev-Web/resolve/main/python/predictor.py
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()), | |
| } | |