Download code/src/stackcraft/clef.py from nima1/stackcraft-clef-flash-lora: direct link, hf CLI and curl.
- Browser
- Download file 11.9 kB
-
https://huggingface.co/nima1/stackcraft-clef-flash-lora/resolve/main/code/src/stackcraft/clef.py
- Command line
-
hf download hf://nima1/stackcraft-clef-flash-lora/code/src/stackcraft/clef.py
-
curl -L -o clef.py https://huggingface.co/nima1/stackcraft-clef-flash-lora/resolve/main/code/src/stackcraft/clef.py
11.9 kB
| """Pinned native Clef integration; importing this module needs no ML packages.""" | |
| from __future__ import annotations | |
| import hashlib | |
| import importlib | |
| import json | |
| import sys | |
| from pathlib import Path | |
| from types import ModuleType | |
| from typing import Any | |
| from stackcraft.players import Decision, Observation, validate_decision | |
| from stackcraft.schema import RULES_VERSION | |
| MODEL_ID = "Cloudflare/clef-flash" | |
| MODEL_REVISION = "17f0b0ad64efb65d273590632833508766b2aae6" | |
| SOURCE_SHA256 = "0e304cf7c6500e8bb59bef7e2afd2c6373f82596dfb3b57d1aa93c175e2dc3a3" | |
| ENCODING_VERSION = "stackcraft-clef-v1" | |
| QUESTION_ID = "placement" | |
| DEFAULT_MAX_LENGTH = 4096 | |
| _INSTRUCTIONS = ( | |
| "Choose the legal placement that maximizes total lines cleared over the game. " | |
| "Avoid holes and high stacks so future pieces can be placed. Only the current " | |
| "piece and exactly one next piece are known. Consider all supplied placements." | |
| ) | |
| def observation_record(observation: Observation) -> dict[str, Any]: | |
| """Build one native choice question from visible state; never include a seed.""" | |
| if observation.rules_version != RULES_VERSION: | |
| raise ValueError("unsupported observation rules version") | |
| actions = observation.legal_actions | |
| if not actions: | |
| raise ValueError("cannot ask Clef to choose with no legal placements") | |
| if len({action.id for action in actions}) != len(actions): | |
| raise ValueError("legal placement IDs must be unique") | |
| return { | |
| "model": MODEL_ID, | |
| "state": { | |
| "encoding_version": ENCODING_VERSION, | |
| "rules_version": RULES_VERSION, | |
| "board_rows": [ | |
| "".join("#" if cell else "." for cell in row) for row in observation.board | |
| ], | |
| "current_piece": observation.current, | |
| "next_piece": observation.next_piece, | |
| "coordinates": ( | |
| "10 columns x=0..9 left to right; 20 rows y=0..19 top to bottom. " | |
| "board_rows are top to bottom; . is empty and # is occupied. " | |
| "Placement cells use absolute [x,y] coordinates." | |
| ), | |
| "rules": ( | |
| "Place the current four-cell piece at one supplied legal landing. " | |
| "A vertical hard drop starts fully inside row 0; no tucks, wall kicks, " | |
| "hold or gravity timer. Full rows clear simultaneously; rows above fall. " | |
| "Score for 1/2/3/4 cleared rows is 100/300/500/800. " | |
| "The next piece becomes current. No legal placement means game over. " | |
| "Future pieces beyond the one preview are unknown." | |
| ), | |
| }, | |
| "questions": { | |
| QUESTION_ID: { | |
| "type": "choice", | |
| "instructions": _INSTRUCTIONS, | |
| "criteria": { | |
| action.id: { | |
| "rotation": action.rotation, | |
| "column": action.x, | |
| "landing_row": action.y, | |
| "cells": [list(cell) for cell in action.cells], | |
| } | |
| for action in actions | |
| }, | |
| } | |
| }, | |
| } | |
| def _render(value: Any) -> str: | |
| return ( | |
| value | |
| if isinstance(value, str) | |
| else json.dumps(value, ensure_ascii=False, separators=(",", ":"), sort_keys=True) | |
| ) | |
| def complete_token_count(tokenizer: Any, native: Any, record: dict[str, Any]) -> int: | |
| """Count exact text segments used by the pinned native encoder, before encoding. | |
| Tokenizing the joined text is NOT equivalent: native encode_record tokenizes | |
| each segment separately. This implementation is coupled to SOURCE_SHA256. | |
| Only the one-question, text-only Stackcraft record schema is supported. | |
| """ | |
| if record.get("images") or record.get("videos"): | |
| raise ValueError("Stackcraft Clef encoding is text-only") | |
| questions = record.get("questions", {}) | |
| if list(questions) != [QUESTION_ID] or questions[QUESTION_ID].get("type") != "choice": | |
| raise ValueError("expected the single Stackcraft placement choice question") | |
| question = questions[QUESTION_ID] | |
| segments = [ | |
| "\n\nSCHEMA FIELDS:\n", | |
| f"\nFIELD 1\nID: {QUESTION_ID}\nTYPE: choice\nINSTRUCTION: ", | |
| _render(question["instructions"]), | |
| "\nALLOWED OPTIONS:\n", | |
| ] | |
| for index, (option_id, description) in enumerate(sorted(question["criteria"].items())): | |
| semantics = {"option_id": option_id} | |
| if description is not None: | |
| semantics["description"] = description | |
| segments.extend((f"OPTION {index + 1}: ", _render(semantics), "\n")) | |
| segments.extend( | |
| ( | |
| "END FIELD\n", | |
| f"<|im_start|>system\n{native.SYSTEM_PROMPT}<|im_end|>\n<|im_start|>user\nSTATE:\n", | |
| "\n<|im_end|>\n<|im_start|>assistant\n<think>\n\n</think>\n\nJOINT SCHEMA DECISIONS:", | |
| _render(record["state"]), | |
| ) | |
| ) | |
| return sum(len(tokenizer(segment, add_special_tokens=False).input_ids) for segment in segments) | |
| def encode_observation( | |
| observation: Observation, tokenizer: Any, native: Any, max_length: int = DEFAULT_MAX_LENGTH | |
| ) -> Any: | |
| """Reject overlong input before native encode_record can truncate the board.""" | |
| if type(max_length) is not int or max_length < 1: | |
| raise ValueError("max_length must be a positive integer") | |
| record = observation_record(observation) | |
| required = complete_token_count(tokenizer, native, record) | |
| if required > max_length: | |
| raise ValueError( | |
| f"complete Stackcraft state and choices require {required} tokens; " | |
| f"max_length={max_length}; refusing to truncate" | |
| ) | |
| encoded = native.encode_record(tokenizer, record, max_length=max_length) | |
| if len(encoded.input_ids) != required: | |
| raise ValueError("native encoding length differs from preflight; source contract changed") | |
| expected_ids = tuple(sorted(action.id for action in observation.legal_actions)) | |
| if ( | |
| len(encoded.questions) != 1 | |
| or encoded.questions[0].question_id != QUESTION_ID | |
| or encoded.questions[0].option_ids != expected_ids | |
| ): | |
| raise ValueError("native encoded option IDs differ from complete legal action set") | |
| return encoded | |
| def import_pinned_source(path: Path, *, trust_pinned_code: bool = False) -> ModuleType: | |
| """Execute only the explicitly trusted, reviewed native source bytes. | |
| A pinned revision limits changes; it does not make Python code a sandbox. | |
| The caller must deliberately accept executing the reviewed upstream module. | |
| """ | |
| if not trust_pinned_code: | |
| raise ValueError("native Clef loading requires trust_pinned_code=True") | |
| source = path.read_bytes() | |
| if hashlib.sha256(source).hexdigest() != SOURCE_SHA256: | |
| raise ValueError("native Clef source SHA256 mismatch; refusing to execute") | |
| name = f"_stackcraft_clef_native_{SOURCE_SHA256}" | |
| if name in sys.modules: | |
| return sys.modules[name] | |
| module = ModuleType(name) | |
| module.__file__ = str(path) | |
| sys.modules[name] = module # dataclasses resolves annotations through this registry. | |
| try: | |
| exec(compile(source, str(path), "exec"), module.__dict__) | |
| except BaseException: | |
| del sys.modules[name] | |
| raise | |
| return module | |
| class ClefPlayer: | |
| """Native single-record inference with unrounded probabilities. | |
| Direct construction also supports an explicitly supplied trained native model; | |
| such callers must set revision to the actual checkpoint identity. The factory | |
| below loads only the pinned, unchanged upstream release. | |
| """ | |
| name = "clef-flash" | |
| def __init__( | |
| self, | |
| model: Any, | |
| processor: Any, | |
| native: Any, | |
| *, | |
| revision: str, | |
| max_length: int = DEFAULT_MAX_LENGTH, | |
| ) -> None: | |
| if type(max_length) is not int or max_length < 1: | |
| raise ValueError("max_length must be a positive integer") | |
| self.model = model.eval() | |
| self.processor = processor | |
| self.native = native | |
| self.revision = revision | |
| self.max_length = max_length | |
| self.last_input_tokens: int | None = None | |
| self.runtime_config: dict[str, Any] = { | |
| "model_id": MODEL_ID, | |
| "revision": revision, | |
| "base_revision": MODEL_REVISION, | |
| "encoding_version": ENCODING_VERSION, | |
| "source_sha256": SOURCE_SHA256, | |
| "max_length": max_length, | |
| "dtype": None, | |
| "device": None, | |
| } | |
| def from_pretrained( | |
| cls, | |
| *, | |
| trust_pinned_code: bool = False, | |
| local_files_only: bool = True, | |
| device: str = "cuda", | |
| max_length: int = DEFAULT_MAX_LENGTH, | |
| ) -> ClefPlayer: | |
| """Load a pinned local snapshot; downloading is separately opt-in. | |
| This loads real 9B weights onto device. It performs no GPU admission or | |
| workload management; callers must establish available memory beforehand. | |
| """ | |
| if not trust_pinned_code: | |
| raise ValueError("native Clef loading requires trust_pinned_code=True") | |
| if type(max_length) is not int or max_length < 1: | |
| raise ValueError("max_length must be a positive integer") | |
| hub = importlib.import_module("huggingface_hub") | |
| snapshot = Path( | |
| hub.snapshot_download( | |
| repo_id=MODEL_ID, | |
| revision=MODEL_REVISION, | |
| local_files_only=local_files_only, | |
| ) | |
| ) | |
| native = import_pinned_source( | |
| snapshot / "joint_schema_model.py", trust_pinned_code=trust_pinned_code | |
| ) | |
| config = json.loads((snapshot / "config.json").read_text()) | |
| context = config["text_config"]["max_position_embeddings"] | |
| if max_length > context: | |
| raise ValueError(f"max_length={max_length} exceeds backbone context={context}") | |
| torch = importlib.import_module("torch") | |
| model, processor = native.load_release_model(snapshot, device=device, dtype=torch.bfloat16) | |
| return cls( | |
| model, | |
| processor, | |
| native, | |
| revision=f"{MODEL_ID}@{MODEL_REVISION}:{ENCODING_VERSION}", | |
| max_length=max_length, | |
| ) | |
| def choose(self, observation: Observation) -> Decision: | |
| tokenizer = self.processor.tokenizer | |
| encoded = encode_observation(observation, tokenizer, self.native, self.max_length) | |
| pad_id = tokenizer.pad_token_id | |
| if pad_id is None: | |
| raise ValueError("native Clef tokenizer has no padding token") | |
| torch = importlib.import_module("torch") | |
| first_parameter = next(self.model.parameters()) | |
| device = first_parameter.device | |
| batch = self.native.collate_records([encoded], pad_id, device) | |
| with torch.inference_mode(): | |
| result = self.model(batch) | |
| if len(result) != 1 or len(result[0]) != 1: | |
| raise ValueError("native Clef must return one batch and one question") | |
| logits = result[0][0] | |
| values = logits.float().softmax(-1).tolist() | |
| ids = encoded.questions[0].option_ids | |
| if len(values) != len(ids): | |
| raise ValueError("native Clef probability count differs from legal choices") | |
| probabilities = dict(zip(ids, values, strict=True)) | |
| # Preserve common engine ordering on exact ties, not native lexical order. | |
| selected = max(observation.legal_actions, key=lambda action: probabilities[action.id]) | |
| decision = Decision(selected.id, probabilities) | |
| validate_decision(decision, observation) | |
| self.last_input_tokens = len(encoded.input_ids) | |
| self.runtime_config.update( | |
| dtype=str(first_parameter.dtype), | |
| device=str(device), | |
| ) | |
| return decision | |