Download code/src/stackcraft/training.py from nima1/stackcraft-clef-flash-lora: direct link, hf CLI and curl.
- Browser
- Download file 13.9 kB
-
https://huggingface.co/nima1/stackcraft-clef-flash-lora/resolve/main/code/src/stackcraft/training.py
- Command line
-
hf download hf://nima1/stackcraft-clef-flash-lora/code/src/stackcraft/training.py
-
curl -L -o training.py https://huggingface.co/nima1/stackcraft-clef-flash-lora/resolve/main/code/src/stackcraft/training.py
13.9 kB
| """Optional native Clef training primitives; imported only by ML workflows.""" | |
| from __future__ import annotations | |
| import hashlib | |
| import json | |
| import math | |
| import re | |
| from pathlib import Path | |
| from typing import Any, Literal | |
| import torch | |
| from safetensors.torch import load_file, save_file | |
| from torch import nn | |
| from torch.nn import functional as functional | |
| from stackcraft.clef import ENCODING_VERSION, MODEL_ID, MODEL_REVISION, SOURCE_SHA256 | |
| FORMAT_VERSION = 1 | |
| HEAD_TYPE = "native-joint-schema-fp32-gathered-rows-v1" | |
| _TARGET = re.compile( | |
| r"^model\.language_model\.layers\.\d+\." | |
| r"(?:self_attn|linear_attn|mlp)\." | |
| r"(?:q_proj|k_proj|v_proj|o_proj|in_proj_qkv|in_proj_z|in_proj_b|in_proj_a|" | |
| r"out_proj|gate_proj|up_proj|down_proj)$" | |
| ) | |
| class GatheredFloat32Embedding: | |
| """The pinned head only indexes this object; never materialize all rows in FP32.""" | |
| def __init__(self, weight: torch.Tensor) -> None: | |
| self.weight = weight | |
| def __getitem__(self, indices: torch.Tensor) -> torch.Tensor: | |
| # Both indexing and casting retain autograd edges when the input requires it. | |
| return self.weight[indices].float() | |
| class FP32DecisionHead(nn.Module): | |
| """Keep the upstream head unchanged while adapting its floating inputs.""" | |
| def __init__(self, native_head: nn.Module) -> None: | |
| super().__init__() | |
| self.native_head = native_head.float() | |
| def forward( | |
| self, | |
| hidden_states: torch.Tensor, | |
| input_ids: torch.Tensor, | |
| attention_mask: torch.Tensor, | |
| records: list[Any], | |
| output_embedding_weight: torch.Tensor, | |
| ) -> list[list[torch.Tensor]]: | |
| # Disable surrounding autocast so trainable head and its operations stay FP32. | |
| with torch.autocast(device_type=hidden_states.device.type, enabled=False): | |
| return self.native_head( | |
| hidden_states.float(), | |
| input_ids, | |
| attention_mask, | |
| records, | |
| GatheredFloat32Embedding(output_embedding_weight), | |
| ) | |
| def lora_target_modules(backbone: nn.Module) -> list[str]: | |
| """Return full text-layer names; suffix-only matching can accidentally train vision.""" | |
| targets = [ | |
| name | |
| for name, module in backbone.named_modules() | |
| if isinstance(module, nn.Linear) and _TARGET.fullmatch(name) | |
| ] | |
| if not targets: | |
| raise ValueError("no supported Qwen3.5 text-layer LoRA targets found") | |
| return sorted(targets) | |
| def prepare_trainable(model: Any, mode: Literal["head", "lora"] = "head", rank: int = 4) -> Any: | |
| """Prepare a pinned, already admitted/loaded native ClefModel in place. | |
| This function loads no model or weights. The caller must supply the pinned | |
| native model, e.g. ClefPlayer.from_pretrained(...).model after memory admission. | |
| Backbone dtype is preserved; the intended release loader supplies BF16. | |
| """ | |
| if mode not in ("head", "lora"): | |
| raise ValueError("training mode must be head or lora") | |
| if type(rank) is not int or rank < 1: | |
| raise ValueError("LoRA rank must be a positive integer") | |
| if isinstance(model.head, FP32DecisionHead) or hasattr(model, "_stackcraft_training"): | |
| raise ValueError("model is already prepared for Stackcraft training") | |
| if hasattr(model.language_model, "peft_config"): | |
| raise ValueError("expected the unchanged backbone, not an existing PEFT model") | |
| model.language_model.requires_grad_(False) | |
| targets: list[str] = [] | |
| if mode == "lora": | |
| from peft import LoraConfig, get_peft_model | |
| targets = lora_target_modules(model.language_model) | |
| model.language_model = get_peft_model( | |
| model.language_model, | |
| LoraConfig( | |
| r=rank, | |
| lora_alpha=2 * rank, | |
| lora_dropout=0.0, | |
| target_modules=targets, | |
| bias="none", | |
| ), | |
| ) | |
| if hasattr(model.language_model, "gradient_checkpointing_enable"): | |
| model.language_model.gradient_checkpointing_enable( | |
| gradient_checkpointing_kwargs={"use_reentrant": False} | |
| ) | |
| model.head = FP32DecisionHead(model.head) | |
| model.head.requires_grad_(True) | |
| model.train() | |
| if mode == "head": | |
| model.language_model.eval() | |
| model._stackcraft_training = { | |
| "format_version": FORMAT_VERSION, | |
| "base_model": MODEL_ID, | |
| "base_revision": MODEL_REVISION, | |
| "native_source_sha256": SOURCE_SHA256, | |
| "encoding_version": ENCODING_VERSION, | |
| "head_type": HEAD_TYPE, | |
| "mode": mode, | |
| "lora": ( | |
| {"rank": rank, "alpha": 2 * rank, "dropout": 0.0, "target_modules": targets} | |
| if mode == "lora" | |
| else None | |
| ), | |
| "loss": {"label_smoothing": 0.05, "brier_weight": 0.1, "brier_reduction": "sum_options"}, | |
| } | |
| return model | |
| def decision_loss( | |
| logits: torch.Tensor, | |
| encoded_record: Any, | |
| target_action_id: str, | |
| *, | |
| label_smoothing: float = 0.05, | |
| brier_weight: float = 0.1, | |
| ) -> torch.Tensor: | |
| """One native choice: smoothed cross entropy plus multiclass Brier sum.""" | |
| if len(encoded_record.questions) != 1: | |
| raise ValueError("training records must contain exactly one choice question") | |
| question = encoded_record.questions[0] | |
| if question.question_type != 1: | |
| raise ValueError("training question must be native choice type 1") | |
| ids = question.option_ids | |
| if tuple(ids) != tuple(sorted(set(ids))): | |
| raise ValueError("encoded option IDs must be unique and lexicographically sorted") | |
| if target_action_id not in ids: | |
| raise ValueError("target action is missing from native encoded option IDs") | |
| if logits.ndim != 1 or logits.numel() != len(ids): | |
| raise ValueError("logits shape does not match the encoded choices") | |
| if ( | |
| not math.isfinite(label_smoothing) | |
| or not 0 <= label_smoothing <= 1 | |
| or not math.isfinite(brier_weight) | |
| or brier_weight < 0 | |
| ): | |
| raise ValueError("invalid label smoothing or Brier weight") | |
| if not torch.isfinite(logits).all(): | |
| raise ValueError("decision logits contain nonfinite values") | |
| values = logits.float().unsqueeze(0) | |
| target = torch.tensor([ids.index(target_action_id)], device=values.device) | |
| cross_entropy = functional.cross_entropy(values, target, label_smoothing=label_smoothing) | |
| one_hot = functional.one_hot(target, num_classes=len(ids)).float() | |
| brier = (values.softmax(-1) - one_hot).square().sum(-1).mean() | |
| return cross_entropy + brier_weight * brier | |
| def parameter_hashes( | |
| model: nn.Module, *, trainable: bool, chunk_elements: int = 1_048_576 | |
| ) -> dict[str, str]: | |
| """Hash selected parameters exactly, moving only bounded chunks to CPU. | |
| Full frozen-backbone hashing is intentionally an explicit before/after audit, | |
| not a training-step operation. Dtype and shape are included in every digest. | |
| """ | |
| if type(chunk_elements) is not int or chunk_elements < 1: | |
| raise ValueError("chunk_elements must be a positive integer") | |
| results = {} | |
| for name, parameter in model.named_parameters(): | |
| if parameter.requires_grad != trainable: | |
| continue | |
| digest = hashlib.sha256() | |
| digest.update(f"{parameter.dtype}:{tuple(parameter.shape)}:".encode()) | |
| flattened = parameter.detach().reshape(-1) | |
| for start in range(0, flattened.numel(), chunk_elements): | |
| chunk = flattened[start : start + chunk_elements].to("cpu").contiguous() | |
| digest.update(chunk.view(torch.uint8).numpy().tobytes()) | |
| results[name] = digest.hexdigest() | |
| return results | |
| def save_checkpoint( | |
| model: Any, path: str | Path, *, extra_metadata: dict[str, Any] | None = None | |
| ) -> dict[str, Any]: | |
| """Save head and optional LoRA separately, never the frozen multi-GB backbone.""" | |
| if not isinstance(model.head, FP32DecisionHead) or not hasattr(model, "_stackcraft_training"): | |
| raise ValueError("model must be prepared before saving a training checkpoint") | |
| destination = Path(path) | |
| destination.mkdir(parents=True, exist_ok=False) | |
| metadata = dict(model._stackcraft_training) | |
| metadata["extra"] = extra_metadata or {} | |
| # Store native head keys, not wrapper-specific state_dict prefixes. | |
| head_state = { | |
| name: tensor.detach().cpu().contiguous() | |
| for name, tensor in model.head.native_head.state_dict().items() | |
| } | |
| save_file(head_state, destination / "joint_head.safetensors") | |
| metadata["head_shapes"] = {name: list(tensor.shape) for name, tensor in head_state.items()} | |
| if metadata["mode"] == "lora": | |
| model.language_model.save_pretrained(destination / "adapter", safe_serialization=True) | |
| adapter_path = destination / "adapter" / "adapter_config.json" | |
| adapter_config = json.loads(adapter_path.read_text()) | |
| adapter_config.update( | |
| base_model_name_or_path=MODEL_ID, | |
| revision=MODEL_REVISION, | |
| target_modules=metadata["lora"]["target_modules"], | |
| ) | |
| adapter_path.write_text(json.dumps(adapter_config, indent=2, sort_keys=True) + "\n") | |
| (destination / "training_config.json").write_text( | |
| json.dumps(metadata, indent=2, sort_keys=True, allow_nan=False) + "\n" | |
| ) | |
| return metadata | |
| def load_checkpoint(model: Any, path: str | Path, *, trainable: bool = False) -> Any: | |
| """Restore onto an unchanged pinned native base; reject incompatible metadata.""" | |
| source = Path(path) | |
| metadata = json.loads((source / "training_config.json").read_text()) | |
| expected = { | |
| "format_version": FORMAT_VERSION, | |
| "base_model": MODEL_ID, | |
| "base_revision": MODEL_REVISION, | |
| "native_source_sha256": SOURCE_SHA256, | |
| "encoding_version": ENCODING_VERSION, | |
| "head_type": HEAD_TYPE, | |
| } | |
| for key, value in expected.items(): | |
| if type(metadata.get(key)) is not type(value) or metadata[key] != value: | |
| raise ValueError(f"checkpoint {key} is incompatible with this pinned native adapter") | |
| mode = metadata.get("mode") | |
| if mode not in ("head", "lora"): | |
| raise ValueError("checkpoint training mode is invalid") | |
| if isinstance(model.head, FP32DecisionHead) or hasattr(model.language_model, "peft_config"): | |
| raise ValueError("checkpoint must load onto an unchanged native base") | |
| head_state = load_file(source / "joint_head.safetensors", device="cpu") | |
| shapes = {name: list(tensor.shape) for name, tensor in head_state.items()} | |
| base_shapes = {name: list(tensor.shape) for name, tensor in model.head.state_dict().items()} | |
| if shapes != metadata.get("head_shapes") or shapes != base_shapes: | |
| raise ValueError("checkpoint head structure differs from metadata or native model") | |
| if any(tensor.dtype != torch.float32 for tensor in head_state.values()): | |
| raise ValueError("checkpoint head tensors must be FP32") | |
| model.language_model.requires_grad_(False) | |
| if mode == "lora": | |
| from peft import LoraConfig, PeftModel | |
| from peft.tuners.tuners_utils import check_target_module_exists | |
| lora = metadata.get("lora") | |
| if not isinstance(lora, dict) or lora.get("target_modules") != lora_target_modules( | |
| model.language_model | |
| ): | |
| raise ValueError("checkpoint LoRA targets differ from the native text backbone") | |
| if ( | |
| type(lora.get("rank")) is not int | |
| or lora["rank"] < 1 | |
| or lora.get("alpha") != 2 * lora["rank"] | |
| or lora.get("dropout") != 0.0 | |
| ): | |
| raise ValueError( | |
| "checkpoint LoRA rank, alpha or dropout violates the training contract" | |
| ) | |
| config = json.loads((source / "adapter" / "adapter_config.json").read_text()) | |
| # PEFT 0.21.2 minimizes >=20 explicit module names to equivalent suffixes. | |
| # Compare their meaning on this exact unchanged backbone, not list spelling. | |
| # Enumerating ALL modules ensures an accidental vision/MTP/lm_head match | |
| # makes the sets unequal and is rejected before installing the adapter. | |
| saved_config = LoraConfig.from_pretrained(str(source / "adapter")) | |
| resolved_targets = sorted( | |
| name | |
| for name, _ in model.language_model.named_modules() | |
| if check_target_module_exists(saved_config, name) | |
| ) | |
| if ( | |
| config.get("r") != lora.get("rank") | |
| or config.get("lora_alpha") != lora.get("alpha") | |
| or config.get("lora_dropout") != lora.get("dropout") | |
| or resolved_targets != lora["target_modules"] | |
| or config.get("bias") != "none" | |
| or config.get("modules_to_save") is not None | |
| or config.get("target_parameters") is not None | |
| ): | |
| raise ValueError("saved adapter configuration differs from checkpoint metadata") | |
| model.language_model = PeftModel.from_pretrained( | |
| model.language_model, source / "adapter", is_trainable=trainable | |
| ) | |
| if trainable and hasattr(model.language_model, "gradient_checkpointing_enable"): | |
| model.language_model.gradient_checkpointing_enable( | |
| gradient_checkpointing_kwargs={"use_reentrant": False} | |
| ) | |
| elif metadata.get("lora") is not None: | |
| raise ValueError("head-only checkpoint must not contain LoRA configuration") | |
| model.head = FP32DecisionHead(model.head) | |
| model.head.native_head.load_state_dict(head_state, strict=True) | |
| model.head.requires_grad_(trainable) | |
| model._stackcraft_training = { | |
| key: value for key, value in metadata.items() if key not in ("extra", "head_shapes") | |
| } | |
| model.train(trainable) | |
| if mode == "head": | |
| model.language_model.eval() | |
| return model | |