Download tests/test_model_training.py from devildasdf/NEXORA: direct link, hf CLI and curl.
- Browser
- Download file 2.84 kB
-
https://huggingface.co/devildasdf/NEXORA/resolve/main/tests/test_model_training.py
- Command line
-
hf download hf://devildasdf/NEXORA/tests/test_model_training.py
-
curl -L -o test_model_training.py https://huggingface.co/devildasdf/NEXORA/resolve/main/tests/test_model_training.py
2.84 kB
| import json | |
| from pathlib import Path | |
| import pytest | |
| import torch | |
| from nexora.model import ModelConfig, NexoraLM | |
| from nexora.tokenizer import ByteTokenizer | |
| from nexora.data import prepare | |
| from nexora.training import train, load_checkpoint | |
| torch.set_num_threads(2) | |
| def test_tokenizer_roundtrip(text): | |
| t = ByteTokenizer() | |
| assert t.decode(t.encode(text, special=True)) == text | |
| def test_invalid_config(): | |
| with pytest.raises(ValueError): | |
| ModelConfig(hidden_size=15) | |
| def test_causal_and_shapes(): | |
| torch.manual_seed(1) | |
| m = NexoraLM(ModelConfig(hidden_size=32, layers=2, heads=4, kv_heads=2, intermediate_size=64, max_context=16)).eval() | |
| x = torch.randint(0, 259, (2, 8)) | |
| y = x.clone() | |
| y[:, 5:] = torch.randint(0, 259, (2, 3)) | |
| a, loss = m(x, x) | |
| b, _ = m(y) | |
| assert a.shape == (2, 8, 259) | |
| torch.testing.assert_close(a[:, :5], b[:, :5], rtol=1e-5, atol=1e-6) | |
| loss.backward() | |
| assert all(p.grad is not None and torch.isfinite(p.grad).all() for p in m.parameters()) | |
| def test_context_guard(): | |
| m = NexoraLM(ModelConfig(max_context=4)) | |
| with pytest.raises(ValueError): | |
| m(torch.zeros((1, 5), dtype=torch.long)) | |
| def test_checkpoint_resume_exact(tmp_path): | |
| cfg = {"model": {"hidden_size": 32, "layers": 1, "heads": 4, "kv_heads": 2, "intermediate_size": 64, "max_context": 32}, | |
| "training": {"steps": 6, "batch_size": 2, "sequence_length": 16, "learning_rate": .001, "seed": 11, "eval_every": 2, "checkpoint_every": 2, "device": "cpu", "threads": 2}} | |
| cp = tmp_path / "config.json" | |
| cp.write_text(json.dumps(cfg)) | |
| prepare([{"id": "1", "text": "An original training document with numbers one two three and code expressions.", "source": "test", "license": "MIT", "domain": "text"}, | |
| {"id": "2", "text": "Validation passages should contain separate statements about model behavior and arithmetic.", "source": "test", "license": "MIT", "domain": "text", "split": "validation"}], tmp_path / "data") | |
| train(cp, tmp_path / "data", tmp_path / "full") | |
| train(cp, tmp_path / "data", tmp_path / "resumed", stop_after=3) | |
| train(cp, tmp_path / "data", tmp_path / "resumed", resume=True) | |
| from safetensors.torch import load_file | |
| a, b = [load_file(str(tmp_path / p / "model.safetensors")) for p in ("full", "resumed")] | |
| assert all(torch.equal(a[k], b[k]) for k in a) | |
| root = tmp_path / "resumed" / "checkpoints" | |
| latest = json.loads((root / "latest.json").read_text()) | |
| with (root / latest["file"]).open("ab") as f: | |
| f.write(b"corrupted") | |
| with pytest.raises(ValueError, match="checksum"): | |
| train(cp, tmp_path / "data", tmp_path / "resumed", resume=True) | |