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)