File size: 2,836 Bytes
12496fc | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 | 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)
@pytest.mark.parametrize("text", ["", "hello", "हिन्दी और Hinglish", "\tdef f():\n return 2\n", "🙂∑α²", '{"x": "\\n"}'])
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)
|