Download source/src/vimeml/training/runtime_v2.py from Voltline/vimeml-tiny-ja-v2.1: direct link, hf CLI and curl.
- Browser
- Download file 5.7 kB
-
https://huggingface.co/Voltline/vimeml-tiny-ja-v2.1/resolve/main/source/src/vimeml/training/runtime_v2.py
- Command line
-
hf download hf://Voltline/vimeml-tiny-ja-v2.1/source/src/vimeml/training/runtime_v2.py
-
curl -L -o runtime_v2.py https://huggingface.co/Voltline/vimeml-tiny-ja-v2.1/resolve/main/source/src/vimeml/training/runtime_v2.py
5.7 kB
| """Optional training compilation and frozen-weight initialization for V2.""" | |
| from types import MethodType | |
| from pathlib import Path | |
| import torch | |
| from torch.nn import functional as F | |
| from vimeml.training.model_v2 import RMSNorm | |
| def rmsnorm_operations(hidden, weight, eps): | |
| normalized = hidden.float() | |
| variance = normalized.pow(2).mean(-1, keepdim=True) | |
| normalized = normalized * torch.rsqrt(variance + eps) | |
| return normalized.to(hidden.dtype) * weight | |
| def fuse_rmsnorm(model): | |
| """Compile only pure norm operations, leaving variable-length logits eager.""" | |
| compiled = torch.compile(rmsnorm_operations, fullgraph=True, dynamic=True) | |
| def forward(module, hidden): | |
| return compiled(hidden, module.weight, module.eps) | |
| count = 0 | |
| for module in model.modules(): | |
| if isinstance(module, RMSNorm): | |
| module.forward = MethodType(forward, module) | |
| count += 1 | |
| if not count: | |
| raise ValueError("RMSNorm fusion requires a V2 model.") | |
| return { | |
| "implementation": "torch.compile RMSNorm only", | |
| "norm_modules": count, | |
| "dynamic_shapes": True, | |
| "state_dict_unchanged": True, | |
| "checkpoint_inference": "Plain primitive RMSNorm; training fusion is optional.", | |
| } | |
| def compile_backbone(model): | |
| """Compile the full hidden-state stack, keeping variable-size token loss eager.""" | |
| def hidden_stack(input_ids): | |
| positions = torch.arange(input_ids.shape[1], device=input_ids.device) | |
| hidden = model.dropout( | |
| model.token_embedding(input_ids) + model.position_embedding(positions) | |
| ) | |
| for block in model.blocks: | |
| hidden = block(hidden) | |
| return model.final_norm(hidden) | |
| compiled = torch.compile(hidden_stack, fullgraph=True, dynamic=True) | |
| def forward(module, input_ids, labels=None): | |
| if ( | |
| input_ids.ndim != 2 | |
| or not 1 <= input_ids.shape[1] <= module.config.context_length | |
| ): | |
| raise ValueError("Expected [batch, time] within the model context.") | |
| hidden = compiled(input_ids) | |
| if labels is None: | |
| return module.lm_head(hidden) | |
| if labels.shape != input_ids.shape: | |
| raise ValueError("Labels must match input shape.") | |
| valid = labels != -100 | |
| logits = module.lm_head(hidden[valid]) | |
| return { | |
| "loss_sum": F.cross_entropy(logits, labels[valid], reduction="sum"), | |
| "token_count": valid.sum(), | |
| } | |
| model.forward = MethodType(forward, model) | |
| return { | |
| "implementation": "torch.compile full hidden stack; eager valid-token head and CE", | |
| "dynamic_shapes": True, | |
| "state_dict_unchanged": True, | |
| "checkpoint_inference": "Plain TinyGPTV2", | |
| } | |
| def initialize_weights(model, checkpoint, signatures, optimizer=None, precision=None): | |
| """Start a separate run; optionally carry compatible AdamW moments forward.""" | |
| path = Path(checkpoint) | |
| saved = torch.load(path, map_location="cpu", weights_only=True) | |
| if ( | |
| saved.get("format") != "vimeml_tiny_gpt_v2" | |
| or saved["model_config"] != model.configuration() | |
| ): | |
| raise ValueError("Initialization requires the same V2 model configuration.") | |
| if any( | |
| saved["signatures"][name] != signatures[name] for name in ("tokens", "windows") | |
| ): | |
| raise ValueError( | |
| "Initialization requires the same tokenizer/token store and window index." | |
| ) | |
| if optimizer is not None: | |
| if saved.get("precision") != precision or not saved.get("optimizer", {}).get( | |
| "state" | |
| ): | |
| raise ValueError( | |
| "Optimizer continuation requires matching precision and nonempty state." | |
| ) | |
| source_groups = saved["optimizer"]["param_groups"] | |
| if len(source_groups) != len(optimizer.param_groups): | |
| raise ValueError("Optimizer parameter groups differ.") | |
| for source, target in zip(source_groups, optimizer.param_groups): | |
| if ( | |
| len(source["params"]) != len(target["params"]) | |
| or source["betas"] != target["betas"] | |
| or source["weight_decay"] != target["weight_decay"] | |
| or source["eps"] != target["eps"] | |
| or source.get("amsgrad", False) != target.get("amsgrad", False) | |
| ): | |
| raise ValueError( | |
| "Optimizer continuation requires compatible AdamW groups." | |
| ) | |
| model.load_state_dict(saved["model"]) | |
| if optimizer is not None: | |
| learning_rates = [group["lr"] for group in optimizer.param_groups] | |
| optimizer.load_state_dict(saved["optimizer"]) | |
| for group, rate in zip(optimizer.param_groups, learning_rates): | |
| group["lr"] = rate | |
| details = { | |
| "checkpoint": str(path.resolve()), | |
| "source_step": saved["step"], | |
| "source_trained_windows": saved["total_windows"], | |
| "source_epoch": saved["epoch"], | |
| "source_batch_cursor": saved["batch_cursor"], | |
| "source_signatures": saved["signatures"], | |
| "file_bytes": path.stat().st_size, | |
| "optimizer": ( | |
| "restored AdamW moments and parameter step counters; new local scheduler and data cursors" | |
| if optimizer is not None | |
| else "fresh AdamW; source optimizer, scheduler and cursors are not restored" | |
| ), | |
| "optimizer_state_count": len(saved["optimizer"]["state"]) | |
| if optimizer is not None | |
| else 0, | |
| "verification": "Model configuration and small dataset manifest signatures; no repeated checkpoint SHA256", | |
| } | |
| del saved | |
| return details | |