Mini-Transformer / tests /units /test_transformer.py
AlaBoussoffara's picture
organized code and set up chainlit for demos
2d52135
Raw
History Blame Contribute Delete
15.4 kB
import pytest
import torch
from mini_transformer.configs import ModelCfg
from mini_transformer.transformer import BasicEncoderDecoderTransformer
from mini_transformer.utils import broadcast_padding_mask
def _dev():
return (
torch.device(f"cuda:{torch.cuda.current_device()}")
if torch.cuda.is_available()
else torch.device("cpu")
)
def _tiny_cfg(**over):
# Small, fast config for tests; zero-length safe via max_seq_len >= 0
base = dict(
name="test",
best_checkpoint_path="/tmp/best.ckpt",
latest_checkpoint_path="/tmp/latest.ckpt",
tokenizer="bpe_8k",
vocab_size=32,
d_model=16,
d_ff=32,
num_heads=4,
num_layers=2,
max_seq_len=32,
pad_id=0,
bos_id=1,
eos_id=2,
dropout_rate=0.1,
)
base.update(over)
return ModelCfg(**base)
def _zeros_mask(batch: int, seq_len: int, device) -> torch.Tensor:
return torch.zeros(batch, seq_len, dtype=torch.bool, device=device)
def _broadcast_mask(mask_2d: torch.Tensor, num_heads: int) -> torch.Tensor:
return broadcast_padding_mask(mask_2d, num_heads)
# -----------------------
# Constructor & config validation
# -----------------------
def test_ctor_cfg_type_and_values():
with pytest.raises(TypeError):
BasicEncoderDecoderTransformer(cfg="not a cfg")
# invalid vocab / dims / indices
bads = [
dict(vocab_size=0),
dict(d_model=0),
dict(max_seq_len=-1),
dict(pad_id=-1),
dict(pad_id=100), # out of range
dict(bos_id=999),
dict(eos_id=999),
]
for upd in bads:
with pytest.raises((TypeError, ValueError)):
BasicEncoderDecoderTransformer(_tiny_cfg(**upd))
# happy path
model = BasicEncoderDecoderTransformer(_tiny_cfg())
# weight tying sanity: same storage
assert model.embed.token_embed.weight.data_ptr() == model.lm_head.fc.weight.data_ptr()
def test_layer_norm_style_auto_selection():
shallow = BasicEncoderDecoderTransformer(_tiny_cfg(num_layers=4))
assert shallow.layer_norm_style == "post"
assert all(not layer.pre_norm for layer in shallow.encoder.layers)
assert all(not layer.pre_norm for layer in shallow.decoder.layers)
deep = BasicEncoderDecoderTransformer(_tiny_cfg(num_layers=6))
assert deep.layer_norm_style == "pre"
assert all(layer.pre_norm for layer in deep.encoder.layers)
assert all(layer.pre_norm for layer in deep.decoder.layers)
def test_layer_norm_style_override_respected():
cfg = _tiny_cfg(num_layers=6, layer_norm_style="post")
model = BasicEncoderDecoderTransformer(cfg)
assert model.layer_norm_style == "post"
assert all(not layer.pre_norm for layer in model.encoder.layers)
assert all(not layer.pre_norm for layer in model.decoder.layers)
# -----------------------
# Encode/Decode/Forward — shapes, dtypes, devices, zero-length
# -----------------------
@pytest.mark.parametrize("Sx,Sy", [(5, 6)])
def test_forward_pipeline_shapes_and_types(Sx, Sy):
device = _dev()
cfg = _tiny_cfg(max_seq_len=16)
model = BasicEncoderDecoderTransformer(cfg).to(device)
B = 2
src = torch.randint(0, cfg.vocab_size, (B, Sx), dtype=torch.long, device=device)
tgt = torch.randint(0, cfg.vocab_size, (B, Sy), dtype=torch.long, device=device)
# padding masks (boolean) shaped [B, S]; zero-length handled via empty tensors
src_mask_2d = _zeros_mask(B, Sx, device)
tgt_mask_2d = _zeros_mask(B, Sy, device)
logits = model(src, tgt, src_mask_2d, tgt_mask_2d)
assert logits.shape == (B, Sy, cfg.vocab_size)
assert logits.device == device
assert logits.dtype == torch.get_default_dtype()
# encode / decode individually
src_mask = _broadcast_mask(src_mask_2d, cfg.num_heads)
tgt_mask = _broadcast_mask(tgt_mask_2d, cfg.num_heads)
mem = model.encode(src, src_mask)
assert mem.shape == (B, Sx, cfg.d_model)
out = model.decode(mem, tgt, src_mask, tgt_mask)
assert out.shape == (B, Sy, cfg.d_model)
def test_forward_dtype_checks_and_shape_errors():
cfg = _tiny_cfg()
model = BasicEncoderDecoderTransformer(cfg)
B, Sx, Sy = 2, 4, 5
src = torch.randint(0, cfg.vocab_size, (B, Sx), dtype=torch.long)
tgt = torch.randint(0, cfg.vocab_size, (B, Sy), dtype=torch.long)
src_mask = _zeros_mask(B, Sx, src.device)
tgt_mask = _zeros_mask(B, Sy, tgt.device)
with pytest.raises(TypeError):
model("not a tensor", tgt, src_mask, tgt_mask)
with pytest.raises(TypeError):
model(src, "not a tensor", src_mask, tgt_mask)
with pytest.raises(ValueError):
bad_src = src.unsqueeze(0)
bad_src_mask = _zeros_mask(bad_src.shape[0], bad_src.shape[1], bad_src.device)
model(bad_src, tgt, bad_src_mask, tgt_mask)
with pytest.raises(ValueError):
bad_tgt = tgt.unsqueeze(-1)
model(src, bad_tgt, src_mask, tgt_mask)
with pytest.raises(TypeError):
model(src.float(), tgt, src_mask, tgt_mask) # wrong dtype
with pytest.raises(TypeError):
model(src, tgt.int(), src_mask, tgt_mask) # wrong dtype
with pytest.raises(TypeError):
model(src, tgt, src_padding_mask="not a tensor", tgt_padding_mask=tgt_mask)
with pytest.raises(TypeError):
model(src, tgt, src_mask, tgt_padding_mask="not a tensor")
def test_encode_decode_errors_and_messages():
cfg = _tiny_cfg()
model = BasicEncoderDecoderTransformer(cfg)
src = torch.randint(0, cfg.vocab_size, (2, 4))
src_mask = _broadcast_mask(_zeros_mask(src.shape[0], src.shape[1], src.device), cfg.num_heads)
with pytest.raises(ValueError):
model.encode(src.unsqueeze(-1), src_mask) # rank 3 not allowed
with pytest.raises(TypeError):
model.encode(src.float(), src_mask) # dtype must be long
mem = torch.randn(2, 4, cfg.d_model)
tgt = torch.randint(0, cfg.vocab_size, (2, 5))
tgt_mask = _broadcast_mask(_zeros_mask(tgt.shape[0], tgt.shape[1], tgt.device), cfg.num_heads)
with pytest.raises(ValueError):
model.decode(mem.unsqueeze(0), tgt, src_mask, tgt_mask) # mem rank 4
with pytest.raises(ValueError):
model.decode(mem, tgt.unsqueeze(-1), src_mask, tgt_mask) # tgt rank 3
with pytest.raises(TypeError):
model.decode(mem, tgt.float(), src_mask, tgt_mask) # tgt dtype must be long
def test_exceeding_max_seq_len_raises():
cfg = _tiny_cfg(max_seq_len=4)
model = BasicEncoderDecoderTransformer(cfg)
# src longer than max -> PositionalEmbedding should raise
src = torch.randint(0, cfg.vocab_size, (2, 5), dtype=torch.long)
src_mask = _broadcast_mask(_zeros_mask(src.shape[0], src.shape[1], src.device), cfg.num_heads)
with pytest.raises(ValueError):
model.encode(src, src_mask)
# -----------------------
# Gradients & weight tying
# -----------------------
def test_backward_through_full_forward_and_tied_weights():
device = _dev()
cfg = _tiny_cfg()
model = BasicEncoderDecoderTransformer(cfg).to(device)
B, Sx, Sy = 2, 6, 5
src = torch.randint(0, cfg.vocab_size, (B, Sx), dtype=torch.long, device=device)
tgt = torch.randint(0, cfg.vocab_size, (B, Sy), dtype=torch.long, device=device)
src_mask = _zeros_mask(B, Sx, device)
tgt_mask = _zeros_mask(B, Sy, device)
logits = model(src, tgt, src_mask, tgt_mask) # (B, Sy, V)
# Dummy labels (language modeling): predict tgt itself; ignore pads via mask
loss = logits.pow(2).mean()
loss.backward()
# Tied weights should have gradient (one shared tensor)
tied_weight = model.lm_head.fc.weight
assert tied_weight.grad is not None
assert torch.isfinite(tied_weight.grad).all()
# -----------------------
# Generate API
# -----------------------
def test_generate_arg_checks_and_output_shapes():
device = _dev()
cfg = _tiny_cfg()
model = BasicEncoderDecoderTransformer(cfg).to(device)
B, Sx = 2, 4
src = torch.randint(0, cfg.vocab_size, (B, Sx), dtype=torch.long, device=device)
src_mask = _zeros_mask(B, Sx, device)
# type/value errors
with pytest.raises(TypeError):
model.generate("not", src_mask)
with pytest.raises(ValueError):
model.generate(src.unsqueeze(-1), src_mask)
with pytest.raises(TypeError):
model.generate(src.float(), src_mask)
with pytest.raises(TypeError):
model.generate(src, "not a tensor")
with pytest.raises(TypeError):
model.generate(src, src_mask, max_new_tokens="10")
with pytest.raises(ValueError):
model.generate(src, src_mask, max_new_tokens=-1)
with pytest.raises(TypeError):
model.generate(src, src_mask, temperature="1.0")
with pytest.raises(ValueError):
model.generate(src, src_mask, temperature=0.0)
with pytest.raises(TypeError):
model.generate(src, src_mask, top_k="5")
with pytest.raises(TypeError):
model.generate(src, src_mask, top_p="0.9")
with pytest.raises(ValueError):
model.generate(src, src_mask, top_p=0.0)
with pytest.raises(ValueError):
model.generate(src, src_mask, top_p=1.1)
with pytest.raises(TypeError):
model.generate(src, src_mask, seed="123")
with pytest.raises(TypeError):
model.generate(src, src_mask, generator="not a generator")
# happy path: shapes & dtypes; content is stochastic so we don't assert exact ids
out = model.generate(src, src_mask, max_new_tokens=7, temperature=1.0, top_k=None, top_p=None)
# Should start with BOS and append new tokens; total length = 1 + new tokens (or earlier if EOS)
assert out.shape[0] == B and out.dim() == 2
assert out.dtype == torch.long and out.device == device
assert out.size(1) >= 1 and out.size(1) <= 1 + 7 # early stop on EOS allowed
# Ensure BOS is at position 0 for every sequence
assert torch.all(out[:, 0] == cfg.bos_id)
def test_generate_zero_new_tokens_returns_only_bos():
device = _dev()
cfg = _tiny_cfg()
model = BasicEncoderDecoderTransformer(cfg).to(device)
src = torch.randint(0, cfg.vocab_size, (2, 3), dtype=torch.long, device=device)
src_mask = _zeros_mask(src.shape[0], src.shape[1], device)
out = model.generate(src, src_mask, max_new_tokens=0)
assert out.shape == (2, 1) # just the BOS token
def test_generate_sampling_seed_reproducible():
device = _dev()
cfg = _tiny_cfg()
model = BasicEncoderDecoderTransformer(cfg).to(device).eval()
B, Sx = 2, 4
src = torch.randint(0, cfg.vocab_size, (B, Sx), dtype=torch.long, device=device)
src_mask = _zeros_mask(B, Sx, device)
out_seed_a = model.generate(
src,
src_mask,
max_new_tokens=5,
do_sample=True,
temperature=1.0,
top_k=None,
top_p=None,
seed=123,
)
out_seed_b = model.generate(
src,
src_mask,
max_new_tokens=5,
do_sample=True,
temperature=1.0,
top_k=None,
top_p=None,
seed=123,
)
assert torch.equal(out_seed_a, out_seed_b)
out_seed_c = model.generate(
src,
src_mask,
max_new_tokens=5,
do_sample=True,
temperature=1.0,
top_k=None,
top_p=None,
seed=124,
)
assert not torch.equal(out_seed_a, out_seed_c)
gen1 = torch.Generator(device=device)
gen1.manual_seed(999)
gen2 = torch.Generator(device=device)
gen2.manual_seed(999)
out_gen_a = model.generate(
src,
src_mask,
max_new_tokens=5,
do_sample=True,
temperature=1.0,
top_k=None,
top_p=None,
generator=gen1,
)
out_gen_b = model.generate(
src,
src_mask,
max_new_tokens=5,
do_sample=True,
temperature=1.0,
top_k=None,
top_p=None,
generator=gen2,
)
assert torch.equal(out_gen_a, out_gen_b)
# -----------------------
# Mask broadcast sanity (boolean masks handled in utils; we just pass through)
# -----------------------
def test_forward_accepts_boolean_padding_masks():
device = _dev()
cfg = _tiny_cfg()
model = BasicEncoderDecoderTransformer(cfg).to(device)
B, Sx, Sy = 2, 5, 6
src = torch.randint(0, cfg.vocab_size, (B, Sx), dtype=torch.long, device=device)
tgt = torch.randint(0, cfg.vocab_size, (B, Sy), dtype=torch.long, device=device)
src_pad_bool = _zeros_mask(B, Sx, device)
tgt_pad_bool = _zeros_mask(B, Sy, device)
logits = model(src, tgt, src_pad_bool, tgt_pad_bool)
assert logits.shape == (B, Sy, cfg.vocab_size)
# -----------------------
# Weight tying and optimizer step
# -----------------------
def test_weight_tying_and_optimizer_step_changes_weights_once():
device = _dev()
cfg = _tiny_cfg()
model = BasicEncoderDecoderTransformer(cfg).to(device)
# Same storage pointer means truly tied
assert model.embed.token_embed.weight.data_ptr() == model.lm_head.fc.weight.data_ptr()
opt = torch.optim.SGD(model.parameters(), lr=0.01)
B, Sx, Sy = 2, 5, 4
src = torch.randint(0, cfg.vocab_size, (B, Sx), dtype=torch.long, device=device)
tgt = torch.randint(0, cfg.vocab_size, (B, Sy), dtype=torch.long, device=device)
src_mask = _zeros_mask(B, Sx, device)
tgt_mask = _zeros_mask(B, Sy, device)
logits = model(src, tgt, src_mask, tgt_mask) # (B, Sy, V)
# Simple loss to get gradients flowing
loss = logits.pow(2).mean()
loss.backward()
# The tied tensor should have a gradient and be finite
tied = model.lm_head.fc.weight
assert tied.grad is not None and torch.isfinite(tied.grad).all()
# Save copy before step and ensure it changes after step (once, not twice)
before = tied.detach().clone()
opt.step()
after = tied.detach().clone()
assert not torch.allclose(before, after)
# Still tied after step (same storage)
assert model.embed.token_embed.weight.data_ptr() == model.lm_head.fc.weight.data_ptr()
# -----------------------
# Autoregressive invariants: causal mask honored
# -----------------------
def test_decode_invariance_to_future_tokens():
"""
Sanity-check causal masking: logits at position t must not depend on tokens at > t.
Construct two target sequences that are identical up to t and differ after; compare
decoder outputs at step t.
"""
device = _dev()
cfg = _tiny_cfg()
model = BasicEncoderDecoderTransformer(cfg).to(device).eval()
B, Sx, Sy = 1, 4, 5
src = torch.randint(0, cfg.vocab_size, (B, Sx), dtype=torch.long, device=device)
src_mask_2d = _zeros_mask(B, Sx, device)
src_mask = _broadcast_mask(src_mask_2d, cfg.num_heads)
mem = model.encode(src, src_mask)
# Construct y1 and y2 identical up to t=3, different afterwards
y1 = torch.randint(0, cfg.vocab_size, (B, Sy), dtype=torch.long, device=device)
y2 = y1.clone()
t = 3
if Sy > t + 1:
y2[:, t + 1 :] = (y1[:, t + 1 :] + 1) % cfg.vocab_size
tgt_mask_2d = _zeros_mask(B, Sy, device)
tgt_mask = _broadcast_mask(tgt_mask_2d, cfg.num_heads)
out1 = model.decode(mem, y1, src_mask, tgt_mask) # (B, Sy, D)
out2 = model.decode(mem, y2, src_mask, tgt_mask)
# Compare hidden states up to and including position t
assert torch.allclose(out1[:, : t + 1], out2[:, : t + 1], atol=1e-5, rtol=1e-5)