Spaces:
Sleeping
Sleeping
| 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 | |
| # ----------------------- | |
| 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) | |