Download code/tests/test_training.py from nima1/stackcraft-clef-flash-lora: direct link, hf CLI and curl.
- Browser
- Download file 15.2 kB
-
https://huggingface.co/nima1/stackcraft-clef-flash-lora/resolve/main/code/tests/test_training.py
- Command line
-
hf download hf://nima1/stackcraft-clef-flash-lora/code/tests/test_training.py
-
curl -L -o test_training.py https://huggingface.co/nima1/stackcraft-clef-flash-lora/resolve/main/code/tests/test_training.py
15.2 kB
| """Tiny native-head training checks, not proof that the 9B GPU probe fits.""" | |
| import copy | |
| import json | |
| from pathlib import Path | |
| from types import SimpleNamespace | |
| import pytest | |
| torch = pytest.importorskip("torch", reason="training tests require the optional ML extra") | |
| pytest.importorskip("peft", reason="training tests require the optional ML extra") | |
| from stackcraft.clef import MODEL_ID, MODEL_REVISION, import_pinned_source # noqa: E402 | |
| from stackcraft.training import ( # noqa: E402 | |
| FP32DecisionHead, | |
| GatheredFloat32Embedding, | |
| decision_loss, | |
| load_checkpoint, | |
| lora_target_modules, | |
| parameter_hashes, | |
| prepare_trainable, | |
| save_checkpoint, | |
| ) | |
| def native(): | |
| from huggingface_hub import hf_hub_download | |
| try: | |
| path = hf_hub_download( | |
| MODEL_ID, "joint_schema_model.py", revision=MODEL_REVISION, local_files_only=True | |
| ) | |
| except OSError: | |
| pytest.skip("pinned native source must be cached for actual-head CPU tests") | |
| return import_pinned_source(Path(path), trust_pinned_code=True) | |
| def record(native): | |
| return native.EncodedRecord( | |
| input_ids=tuple(range(16)), | |
| questions=( | |
| native.EncodedQuestion( | |
| question_id="placement", | |
| question_type=1, | |
| question_span=(1, 3), | |
| option_spans=((4, 6), (7, 9), (10, 12)), | |
| option_ids=("r0x0", "r0x1", "r1x0"), | |
| ), | |
| ), | |
| record_id="tiny-test", | |
| ) | |
| class TinyText(torch.nn.Module): | |
| def __init__(self): | |
| super().__init__() | |
| self.embed_tokens = torch.nn.Embedding(32, 8) | |
| layer = torch.nn.Module() | |
| layer.self_attn = torch.nn.Module() | |
| layer.self_attn.q_proj = torch.nn.Linear(8, 8, bias=False) | |
| layer.linear_attn = torch.nn.Module() | |
| layer.linear_attn.in_proj_qkv = torch.nn.Linear(8, 8, bias=False) | |
| layer.linear_attn.in_proj_z = torch.nn.Linear(8, 8, bias=False) | |
| layer.mlp = torch.nn.Module() | |
| layer.mlp.down_proj = torch.nn.Linear(8, 8, bias=False) | |
| self.layers = torch.nn.ModuleList([layer]) | |
| def forward(self, input_ids, **kwargs): | |
| hidden = self.embed_tokens(input_ids) | |
| for layer in self.layers: | |
| hidden = hidden + layer.mlp.down_proj( | |
| ( | |
| layer.self_attn.q_proj(hidden) | |
| + layer.linear_attn.in_proj_qkv(hidden) | |
| + layer.linear_attn.in_proj_z(hidden) | |
| ).tanh() | |
| ) | |
| return SimpleNamespace(last_hidden_state=hidden) | |
| class TinyBackbone(torch.nn.Module): | |
| def __init__(self): | |
| super().__init__() | |
| self.config = {"model_type": "custom", "_name_or_path": "tiny-offline-test"} | |
| self.model = torch.nn.Module() | |
| self.model.language_model = TinyText() | |
| self.model.visual = torch.nn.Module() | |
| self.model.visual.q_proj = torch.nn.Linear(8, 8) | |
| self.mtp = torch.nn.Module() | |
| self.mtp.q_proj = torch.nn.Linear(8, 8) | |
| self.lm_head = torch.nn.Linear(8, 32, bias=False) | |
| def get_output_embeddings(self): | |
| return self.lm_head | |
| def tiny_model(native): | |
| torch.manual_seed(42) | |
| head = native.JointSchemaHead( | |
| hidden_size=8, width=8, routing_layers=1, layers=1, heads=2, feedforward=16 | |
| ) | |
| return native.ClefModel(TinyBackbone().to(torch.bfloat16), head.to(torch.bfloat16)) | |
| def test_gather_cast_matches_full_cast_and_preserves_selected_gradients(): | |
| weight = torch.randn(32, 8, dtype=torch.bfloat16, requires_grad=True) | |
| selected = torch.tensor([2, 7, 2]) | |
| output = GatheredFloat32Embedding(weight)[selected] | |
| assert torch.equal(output, weight.float()[selected]) | |
| output.sum().backward() | |
| assert torch.all(weight.grad[2] == 2) | |
| assert torch.all(weight.grad[7] == 1) | |
| assert torch.count_nonzero(weight.grad).item() == 16 | |
| def test_actual_native_head_float32_wrapper_matches_reference_and_gradients(native): | |
| model = tiny_model(native) | |
| head = model.head.float().eval() | |
| wrapper = FP32DecisionHead(copy.deepcopy(head)).eval() | |
| encoded = record(native) | |
| batch = native.collate_records([encoded], 0, torch.device("cpu")) | |
| hidden = torch.randn(1, 16, 8, dtype=torch.bfloat16, requires_grad=True) | |
| reference_hidden = hidden.detach().clone().requires_grad_(True) | |
| weight = torch.randn(32, 8, dtype=torch.bfloat16, requires_grad=True) | |
| reference_weight = weight.detach().clone().requires_grad_(True) | |
| wrapped = wrapper(hidden, batch["input_ids"], batch["attention_mask"], [encoded], weight)[0][0] | |
| reference = head( | |
| reference_hidden.float(), | |
| batch["input_ids"], | |
| batch["attention_mask"], | |
| [encoded], | |
| reference_weight.float(), | |
| )[0][0] | |
| assert wrapped.dtype == torch.float32 | |
| torch.testing.assert_close(wrapped, reference, rtol=0, atol=0) | |
| wrapped.square().sum().backward() | |
| reference.square().sum().backward() | |
| torch.testing.assert_close(hidden.grad, reference_hidden.grad, rtol=0, atol=0) | |
| torch.testing.assert_close(weight.grad, reference_weight.grad, rtol=0, atol=0) | |
| assert torch.count_nonzero(hidden.grad) > 0 | |
| for (_, actual), (_, expected) in zip( | |
| wrapper.native_head.named_parameters(), head.named_parameters(), strict=True | |
| ): | |
| torch.testing.assert_close(actual.grad, expected.grad, rtol=0, atol=0) | |
| def test_lora_targets_use_full_names_and_exclude_vision_output_and_mtp(native): | |
| targets = lora_target_modules(tiny_model(native).language_model) | |
| assert targets == [ | |
| "model.language_model.layers.0.linear_attn.in_proj_qkv", | |
| "model.language_model.layers.0.linear_attn.in_proj_z", | |
| "model.language_model.layers.0.mlp.down_proj", | |
| "model.language_model.layers.0.self_attn.q_proj", | |
| ] | |
| def test_real_native_forward_gradients_changed_intended_bytes_and_reload(native, tmp_path, mode): | |
| model = prepare_trainable(tiny_model(native), mode=mode) | |
| encoded = record(native) | |
| batch = native.collate_records([encoded], 0, torch.device("cpu")) | |
| assert isinstance(model.head, FP32DecisionHead) | |
| assert all(parameter.dtype == torch.float32 for parameter in model.head.parameters()) | |
| assert all( | |
| parameter.dtype == torch.bfloat16 | |
| for name, parameter in model.language_model.named_parameters() | |
| if "lora_" not in name | |
| ) | |
| before = parameter_hashes(model, trainable=True, chunk_elements=7) | |
| frozen = parameter_hashes(model, trainable=False, chunk_elements=7) | |
| optimizer = torch.optim.AdamW((p for p in model.parameters() if p.requires_grad), lr=0.002) | |
| for _ in range(2): | |
| optimizer.zero_grad(set_to_none=True) | |
| logits = model(batch)[0][0] | |
| loss = decision_loss(logits, encoded, "r0x1") | |
| assert torch.isfinite(loss) | |
| loss.backward() | |
| grads = [p.grad for p in model.head.parameters() if p.grad is not None] | |
| assert all(torch.isfinite(grad).all() for grad in grads) | |
| assert any(torch.count_nonzero(grad) > 0 for grad in grads) | |
| assert all(p.grad is None for p in model.parameters() if not p.requires_grad) | |
| optimizer.step() | |
| after = parameter_hashes(model, trainable=True, chunk_elements=7) | |
| assert any(before[name] != digest for name, digest in after.items() if name.startswith("head.")) | |
| if mode == "lora": | |
| assert any(before[name] != digest for name, digest in after.items() if "lora_" in name) | |
| assert parameter_hashes(model, trainable=False, chunk_elements=7) == frozen | |
| model.eval() | |
| with torch.no_grad(): | |
| expected = model(batch)[0][0].softmax(-1) | |
| destination = tmp_path / mode | |
| metadata = save_checkpoint(model, destination, extra_metadata={"scope": "tiny-cpu-plumbing"}) | |
| assert metadata["mode"] == mode | |
| assert (destination / "joint_head.safetensors").exists() | |
| assert (destination / "adapter").exists() == (mode == "lora") | |
| if mode == "lora": | |
| adapter_config = json.loads((destination / "adapter" / "adapter_config.json").read_text()) | |
| assert adapter_config["base_model_name_or_path"] == MODEL_ID | |
| assert adapter_config["revision"] == MODEL_REVISION | |
| assert adapter_config["target_modules"] == metadata["lora"]["target_modules"] | |
| restored = load_checkpoint(tiny_model(native), destination) | |
| assert not any(parameter.requires_grad for parameter in restored.parameters()) | |
| with torch.no_grad(): | |
| actual = restored(batch)[0][0].softmax(-1) | |
| torch.testing.assert_close(actual, expected, rtol=0, atol=0) | |
| with pytest.raises(FileExistsError): | |
| save_checkpoint(model, destination) | |
| def test_loss_uses_sorted_choice_index_and_declared_formula(native): | |
| encoded = record(native) | |
| logits = torch.tensor([0.2, -0.4, 1.2], requires_grad=True) | |
| actual = decision_loss(logits, encoded, "r0x1") | |
| ce = torch.nn.functional.cross_entropy( | |
| logits.unsqueeze(0), torch.tensor([1]), label_smoothing=0.05 | |
| ) | |
| brier = ((logits.softmax(-1) - torch.tensor([0.0, 1.0, 0.0])) ** 2).sum() | |
| torch.testing.assert_close(actual, ce + 0.1 * brier) | |
| actual.backward() | |
| assert torch.count_nonzero(logits.grad) > 0 | |
| with pytest.raises(ValueError, match="missing"): | |
| decision_loss(logits, encoded, "not-an-action") | |
| with pytest.raises(ValueError, match="nonfinite"): | |
| decision_loss(torch.tensor([0.0, float("nan"), 1.0]), encoded, "r0x1") | |
| def test_checkpoint_rejects_wrong_contract(native, tmp_path, key, value): | |
| model = prepare_trainable(tiny_model(native)) | |
| checkpoint = tmp_path / "checkpoint" | |
| metadata = save_checkpoint(model, checkpoint) | |
| metadata[key] = value | |
| (checkpoint / "training_config.json").write_text(json.dumps(metadata)) | |
| with pytest.raises(ValueError, match=key): | |
| load_checkpoint(tiny_model(native), checkpoint) | |
| def test_parameter_hashes_are_chunk_independent_and_track_dtype_and_values(): | |
| model = torch.nn.Linear(4, 3).to(torch.bfloat16) | |
| first = parameter_hashes(model, trainable=True, chunk_elements=1) | |
| assert first == parameter_hashes(model, trainable=True, chunk_elements=10) | |
| with torch.no_grad(): | |
| model.weight[0, 0] += 1 | |
| second = parameter_hashes(model, trainable=True) | |
| assert first["weight"] != second["weight"] | |
| assert first["bias"] == second["bias"] | |
| assert first["bias"] != parameter_hashes(model.float(), trainable=True)["bias"] | |
| def test_actual_tiny_qwen_hybrid_checkpointing_keeps_lora_gradients(native): | |
| from transformers import Qwen3_5Config, Qwen3_5ForConditionalGeneration | |
| config = Qwen3_5Config( | |
| text_config={ | |
| "vocab_size": 32, | |
| "hidden_size": 8, | |
| "intermediate_size": 16, | |
| "num_hidden_layers": 2, | |
| "num_attention_heads": 1, | |
| "num_key_value_heads": 1, | |
| "head_dim": 8, | |
| "layer_types": ["linear_attention", "full_attention"], | |
| "linear_key_head_dim": 8, | |
| "linear_value_head_dim": 8, | |
| "linear_num_key_heads": 1, | |
| "linear_num_value_heads": 1, | |
| "max_position_embeddings": 64, | |
| "use_cache": False, | |
| "rope_parameters": { | |
| "rope_type": "default", | |
| "rope_theta": 10000, | |
| "partial_rotary_factor": 1.0, | |
| "mrope_section": [1, 1, 2], | |
| }, | |
| }, | |
| vision_config={ | |
| "depth": 1, | |
| "hidden_size": 8, | |
| "intermediate_size": 16, | |
| "num_heads": 1, | |
| "patch_size": 2, | |
| "spatial_merge_size": 1, | |
| "temporal_patch_size": 1, | |
| "out_hidden_size": 8, | |
| "num_position_embeddings": 16, | |
| }, | |
| ) | |
| backbone = Qwen3_5ForConditionalGeneration(config).to(torch.bfloat16) | |
| targets = lora_target_modules(backbone) | |
| assert len(targets) == 15 # 5 linear-attention + 4 full-attention + 6 MLP projections. | |
| assert sum("linear_attn" in name for name in targets) == 5 | |
| assert all(name.startswith("model.language_model.layers.") for name in targets) | |
| head = tiny_model(native).head | |
| model = prepare_trainable(native.ClefModel(backbone, head), mode="lora") | |
| assert model.language_model.get_base_model().is_gradient_checkpointing | |
| encoded = record(native) | |
| batch = native.collate_records([encoded], 0, torch.device("cpu")) | |
| loss = decision_loss(model(batch)[0][0], encoded, "r0x1") | |
| assert torch.isfinite(loss) | |
| loss.backward() | |
| for part in ("linear_attn", "self_attn", "mlp"): | |
| assert any( | |
| parameter.grad is not None and torch.count_nonzero(parameter.grad) > 0 | |
| for name, parameter in model.named_parameters() | |
| if "lora_" in name and part in name | |
| ) | |
| def test_peft_minimized_suffixes_reload_but_extra_targets_fail(native, tmp_path): | |
| def larger_tiny_model(): | |
| model = tiny_model(native) | |
| text = model.language_model.model.language_model | |
| text.layers = torch.nn.ModuleList([copy.deepcopy(text.layers[0]) for _ in range(32)]) | |
| return model | |
| model = prepare_trainable(larger_tiny_model(), mode="lora") | |
| compressed = sorted(model.language_model.peft_config["default"].target_modules) | |
| intended = model._stackcraft_training["lora"]["target_modules"] | |
| assert len(intended) == 128 | |
| assert len(compressed) < len(intended) # Exercise actual PEFT >=20 optimization. | |
| checkpoint = tmp_path / "compressed" | |
| save_checkpoint(model, checkpoint) | |
| config_path = checkpoint / "adapter" / "adapter_config.json" | |
| config = json.loads(config_path.read_text()) | |
| assert config["target_modules"] == intended # New saves are explicit and canonical. | |
| config["target_modules"] = compressed # Recreate the immutable older probe format. | |
| config["base_model_name_or_path"] = "/old/local/cache/snapshot" | |
| config["revision"] = None | |
| config_path.write_text(json.dumps(config)) | |
| restored = load_checkpoint(larger_tiny_model(), checkpoint) | |
| original_adapters = { | |
| name: parameter.detach() for name, parameter in model.named_parameters() if "lora_" in name | |
| } | |
| restored_adapters = { | |
| name: parameter.detach() | |
| for name, parameter in restored.named_parameters() | |
| if "lora_" in name | |
| } | |
| assert original_adapters.keys() == restored_adapters.keys() | |
| for name, expected in original_adapters.items(): | |
| torch.testing.assert_close(restored_adapters[name], expected, rtol=0, atol=0) | |
| # Force an otherwise-valid suffix to also select the output layer. | |
| config["target_modules"] = compressed + ["lm_head"] | |
| config_path.write_text(json.dumps(config)) | |
| with pytest.raises(ValueError, match="adapter configuration"): | |
| load_checkpoint(larger_tiny_model(), checkpoint) | |