NEXORA / tests /test_model_training.py
devildasdf's picture
Release validated NEXORA research prototype, tiny weights and evidence
12496fc verified
Raw History Blame Contribute Delete
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)
@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)