| """ |
| Frox AI Morph 1.1 β Core Model Tests |
| Run with: pytest tests/test_model.py -v |
| |
| These use the `tiny_text_config` / `tiny_config` fixtures (conftest.py) |
| so the whole file runs in a few seconds on CPU with no GPU required. |
| """ |
| from __future__ import annotations |
|
|
| import copy |
| import math |
|
|
| import pytest |
| import torch |
|
|
| from config.model_config import MorphTextConfig, MorphConfig |
| from model.architecture.morph_model import MorphForCausalLM |
| from model.attention.gqa import MorphAttention, MorphDecoderLayer, MorphRMSNorm |
|
|
|
|
| |
|
|
| class TestConfig: |
| def test_scale_configs_are_internally_consistent(self): |
| """hidden_size must equal num_heads * head_dim at every scale.""" |
| for factory in (MorphConfig.scale_1_5b, MorphConfig.scale_3b, MorphConfig.scale_8b): |
| cfg = factory() |
| t = cfg.text |
| assert t.hidden_size == t.num_attention_heads * t.head_dim, ( |
| f"{factory.__name__}: hidden_size={t.hidden_size} != " |
| f"heads({t.num_attention_heads}) * head_dim({t.head_dim})" |
| ) |
| assert t.num_attention_heads % t.num_key_value_heads == 0 |
|
|
| def test_bad_gqa_ratio_raises(self): |
| """num_attention_heads must be divisible by num_key_value_heads.""" |
| with pytest.raises(AssertionError): |
| MorphTextConfig(num_attention_heads=8, num_key_value_heads=3, head_dim=8, hidden_size=64) |
|
|
| def test_odd_head_dim_raises(self): |
| """RoPE requires an even head_dim (rotate_half splits it in two).""" |
| with pytest.raises(AssertionError): |
| MorphTextConfig(num_attention_heads=8, num_key_value_heads=8, head_dim=7, hidden_size=56) |
|
|
| def test_scaled_init_std_shrinks_with_depth(self): |
| shallow = MorphTextConfig(num_hidden_layers=4) |
| deep = MorphTextConfig(num_hidden_layers=64) |
| assert deep.scaled_init_std < shallow.scaled_init_std |
|
|
|
|
| |
|
|
| class TestConstruction: |
| def test_builds_without_error(self, tiny_text_config): |
| model = MorphForCausalLM(tiny_text_config) |
| assert model is not None |
|
|
| def test_param_count_reasonable(self, tiny_text_config): |
| model = MorphForCausalLM(tiny_text_config) |
| params = model.param_count() |
| assert params["total"] > 0 |
| assert params["trainable"] == params["total"] |
|
|
| def test_tied_embeddings_share_storage(self, tiny_text_config): |
| tiny_text_config.tie_word_embeddings = True |
| model = MorphForCausalLM(tiny_text_config) |
| assert model.lm_head.weight.data_ptr() == model.model.embed_tokens.weight.data_ptr() |
|
|
| def test_untied_embeddings_are_independent(self, tiny_text_config): |
| tiny_text_config.tie_word_embeddings = False |
| model = MorphForCausalLM(tiny_text_config) |
| assert model.lm_head.weight.data_ptr() != model.model.embed_tokens.weight.data_ptr() |
|
|
|
|
| |
|
|
| class TestForwardPass: |
| def test_output_shape(self, tiny_text_config, device): |
| model = MorphForCausalLM(tiny_text_config).to(device).eval() |
| input_ids = torch.randint(0, tiny_text_config.vocab_size, (2, 20), device=device) |
|
|
| with torch.no_grad(): |
| out = model(input_ids=input_ids) |
|
|
| assert out.logits.shape == (2, 20, tiny_text_config.total_vocab_size) |
|
|
| def test_no_nans_or_infs(self, tiny_text_config, device): |
| model = MorphForCausalLM(tiny_text_config).to(device).eval() |
| input_ids = torch.randint(0, tiny_text_config.vocab_size, (2, 20), device=device) |
|
|
| with torch.no_grad(): |
| out = model(input_ids=input_ids) |
|
|
| assert not torch.isnan(out.logits).any(), "NaN in logits" |
| assert not torch.isinf(out.logits).any(), "Inf in logits" |
|
|
| def test_loss_computed_when_labels_given(self, tiny_text_config, device): |
| model = MorphForCausalLM(tiny_text_config).to(device).eval() |
| input_ids = torch.randint(0, tiny_text_config.vocab_size, (2, 20), device=device) |
|
|
| with torch.no_grad(): |
| out = model(input_ids=input_ids, labels=input_ids) |
|
|
| assert out.loss is not None |
| assert out.loss.item() > 0 |
| assert not torch.isnan(out.loss) |
|
|
| def test_ignore_index_excludes_masked_labels(self, tiny_text_config, device): |
| """Labels set to -100 must not contribute to the loss.""" |
| model = MorphForCausalLM(tiny_text_config).to(device).eval() |
| input_ids = torch.randint(0, tiny_text_config.vocab_size, (1, 10), device=device) |
|
|
| labels_all = input_ids.clone() |
| labels_half_masked = input_ids.clone() |
| labels_half_masked[:, :5] = -100 |
|
|
| torch.manual_seed(0) |
| with torch.no_grad(): |
| loss_all = model(input_ids=input_ids, labels=labels_all).loss |
| loss_masked = model(input_ids=input_ids, labels=labels_half_masked).loss |
|
|
| |
| assert loss_all.item() != loss_masked.item() |
|
|
| def test_batch_independence(self, tiny_text_config, device): |
| """Padding/content in one batch row must not leak into another (no cross-attention across batch).""" |
| model = MorphForCausalLM(tiny_text_config).to(device).eval() |
| torch.manual_seed(0) |
|
|
| row_a = torch.randint(0, tiny_text_config.vocab_size, (1, 12), device=device) |
| row_b = torch.randint(0, tiny_text_config.vocab_size, (1, 12), device=device) |
| batched = torch.cat([row_a, row_b], dim=0) |
|
|
| with torch.no_grad(): |
| out_batched = model(input_ids=batched).logits |
| out_a_alone = model(input_ids=row_a).logits |
| out_b_alone = model(input_ids=row_b).logits |
|
|
| assert torch.allclose(out_batched[0], out_a_alone[0], atol=1e-4) |
| assert torch.allclose(out_batched[1], out_b_alone[0], atol=1e-4) |
|
|
|
|
| |
|
|
| class TestCausalMasking: |
| """ |
| The single most important correctness property of a decoder-only |
| LM: token i's output must depend only on tokens [0, i], never on |
| anything after it. We verify this numerically rather than just |
| trusting the mask-construction code, by perturbing a future token |
| and checking earlier outputs are bit-for-bit unchanged. |
| """ |
|
|
| def test_future_tokens_do_not_affect_past_outputs(self, tiny_text_config, device): |
| model = MorphForCausalLM(tiny_text_config).to(device).eval() |
| torch.manual_seed(0) |
|
|
| seq_len = 16 |
| input_ids = torch.randint(0, tiny_text_config.vocab_size, (1, seq_len), device=device) |
|
|
| with torch.no_grad(): |
| logits_original = model(input_ids=input_ids).logits |
|
|
| |
| perturbed = input_ids.clone() |
| perturbed[0, -1] = (perturbed[0, -1] + 1) % tiny_text_config.vocab_size |
|
|
| with torch.no_grad(): |
| logits_perturbed = model(input_ids=perturbed).logits |
|
|
| |
| assert torch.allclose( |
| logits_original[:, :-1, :], logits_perturbed[:, :-1, :], atol=1e-5 |
| ), "Changing the last token altered earlier positions' logits β causal mask is leaking future info!" |
|
|
| |
| assert not torch.allclose( |
| logits_original[:, -1, :], logits_perturbed[:, -1, :], atol=1e-5 |
| ), "Last-position logits didn't change at all β suspicious (model may be ignoring input)." |
|
|
| def test_sliding_window_layers_stay_causal(self, tiny_text_config, device): |
| """Same test, but specifically exercises the chunked sliding-window fallback path |
| (sliding_window_size=16 in the fixture, forcing S > window for longer sequences).""" |
| tiny_text_config.sliding_window_size = 4 |
| model = MorphForCausalLM(tiny_text_config).to(device).eval() |
| torch.manual_seed(0) |
|
|
| seq_len = 20 |
| input_ids = torch.randint(0, tiny_text_config.vocab_size, (1, seq_len), device=device) |
|
|
| with torch.no_grad(): |
| logits_original = model(input_ids=input_ids).logits |
|
|
| perturbed = input_ids.clone() |
| perturbed[0, 10] = (perturbed[0, 10] + 1) % tiny_text_config.vocab_size |
|
|
| with torch.no_grad(): |
| logits_perturbed = model(input_ids=perturbed).logits |
|
|
| |
| assert torch.allclose( |
| logits_original[:, :10, :], logits_perturbed[:, :10, :], atol=1e-5 |
| ), "Sliding-window fallback attention leaked future information into earlier positions!" |
|
|
|
|
| |
|
|
| class TestKVCacheEquivalence: |
| """ |
| The output of generating token-by-token with a KV cache must match |
| a full forward pass over the whole sequence at once. This is the |
| property that makes caching a pure speed optimization rather than |
| a behavior change. |
| """ |
|
|
| def test_incremental_matches_full_forward(self, tiny_text_config, device): |
| |
| |
| |
| |
| tiny_text_config.use_sliding_window = False |
| model = MorphForCausalLM(tiny_text_config).to(device).eval() |
| torch.manual_seed(0) |
|
|
| seq_len = 10 |
| input_ids = torch.randint(0, tiny_text_config.vocab_size, (1, seq_len), device=device) |
|
|
| with torch.no_grad(): |
| full_logits = model(input_ids=input_ids).logits |
|
|
| |
| past_key_values = None |
| incremental_logits = [] |
| for t in range(seq_len): |
| out = model( |
| input_ids=input_ids[:, t:t+1], |
| past_key_values=past_key_values, |
| use_cache=True, |
| ) |
| incremental_logits.append(out.logits) |
| past_key_values = out.past_key_values |
|
|
| incremental_logits = torch.cat(incremental_logits, dim=1) |
|
|
| assert torch.allclose(full_logits, incremental_logits, atol=1e-3), ( |
| "Incremental (KV-cached) generation diverged from a full forward " |
| "pass over the same sequence β the cache is not behavior-preserving." |
| ) |
|
|
|
|
| |
|
|
| class TestGeneration: |
| def test_generate_produces_expected_length(self, tiny_text_config, device): |
| model = MorphForCausalLM(tiny_text_config).to(device).eval() |
| input_ids = torch.randint(0, tiny_text_config.vocab_size, (1, 5), device=device) |
|
|
| out = model.generate( |
| input_ids=input_ids, max_new_tokens=10, do_sample=False, |
| eos_token_id=-1, |
| ) |
| assert out.shape[1] == 5 + 10 |
|
|
| def test_generate_stops_at_eos(self, tiny_text_config, device): |
| """Force EOS to be the only high-probability token and confirm generation halts early.""" |
| untied = copy.deepcopy(tiny_text_config) |
| untied.tie_word_embeddings = False |
| model = MorphForCausalLM(untied).to(device).eval() |
|
|
| with torch.no_grad(): |
| model.lm_head.weight.data.zero_() |
| model.lm_head.weight.data[untied.eos_token_id] += 100.0 |
|
|
| input_ids = torch.randint(0, untied.vocab_size, (1, 3), device=device) |
| out = model.generate( |
| input_ids=input_ids, max_new_tokens=20, do_sample=False, |
| eos_token_id=untied.eos_token_id, |
| ) |
| assert out.shape[1] < 3 + 20, "Generation ran the full budget despite EOS being forced" |
| assert out[0, -1].item() == untied.eos_token_id |
|
|
| def test_stream_callback_fires_once_per_token(self, tiny_text_config, device): |
| model = MorphForCausalLM(tiny_text_config).to(device).eval() |
| input_ids = torch.randint(0, tiny_text_config.vocab_size, (1, 3), device=device) |
|
|
| calls = [] |
| model.generate( |
| input_ids=input_ids, max_new_tokens=5, do_sample=False, |
| eos_token_id=-1, stream_callback=lambda b, tok, step: calls.append((b, tok, step)), |
| ) |
| assert len(calls) == 5 |
|
|
|
|
| |
|
|
| class TestSaveLoad: |
| def test_save_and_load_preserves_outputs(self, tiny_text_config, device, tmp_path): |
| model = MorphForCausalLM(tiny_text_config).to(device).eval() |
| input_ids = torch.randint(0, tiny_text_config.vocab_size, (1, 10), device=device) |
|
|
| with torch.no_grad(): |
| logits_before = model(input_ids=input_ids).logits |
|
|
| save_path = tmp_path / "model_ckpt" |
| model.save(str(save_path)) |
|
|
| loaded = MorphForCausalLM.from_saved(str(save_path), device=str(device)).eval() |
| with torch.no_grad(): |
| logits_after = loaded(input_ids=input_ids).logits |
|
|
| assert torch.allclose(logits_before, logits_after, atol=1e-5), ( |
| "Loaded model produces different outputs than the original β " |
| "save/load is not lossless." |
| ) |
|
|