SenlightAI

SenlightAI tests test_config.py

5.7 KB 171 lines Raw ↗ GitHub ↗

"""Config validation and preset parameter counts.The preset names carry parameter counts. If the analytic breakdown and the real module everdisagree, one of them has a bug, so this file asserts they agree and that the name is honestabout what it builds."""from __future__ import annotationsimport reimport pytestfrom elafry.config import (    AffectConfig,    ModelConfig,get_preset,    list_presets,    param_count,)from elafry.models.elafry import Elafry, param_breakdowndef _label_in_millions(name: str):    """Pull the parameter count out of the preset name, in millions.    ``elafry-8b`` -> 8000.0, ``elafry-316m`` -> 316.0. Returns None for names that carry no    count, such as ``elafry-tiny``.    """    m = re.search(r"(\d+(?:\.\d+)?)([mb])$", name)    if not m:        return None    value, unit = float(m.group(1)), m.group(2)    return value * 1000 if unit == "b" else valuedef _total_millions(name: str) -> float:    preset = get_preset(name)    return sum(param_breakdown(preset["model"], preset["affect"]).values()) / 1e6def test_analytic_breakdown_matches_real_module():    """The analytic count must equal ``sum(p.numel())``.    This is the whole reason for keeping an analytic count: it can be checked without    allocating the model, so the check scales to 8B.    """    checked = 0    for name in list_presets():        preset = get_preset(name)        cfg: ModelConfig = preset["model"]        affect: AffectConfig = preset["affect"]        analytic = sum(param_breakdown(cfg, affect).values())        # Skip anything that would need gigabytes just to count its parameters.        if analytic > 400e6:            continue        model = Elafry(cfg, affect)        real = sum(p.numel() for p in model.parameters())        assert analytic == real, (            f"{name}: analytic {analytic:,} != real {real:,}; "            "param_breakdown has drifted from the module construction"        )        checked += 1    assert checked >= 3, f"only checked {checked} presets against the real module"def test_preset_names_match_actual_sizes():    """A preset called 8b should be about 8b.    The breakdown includes the affective subsystem, which every preset switches on, so the    real number sits a little above the bare transformer count. That is the honest number.    """    for name in list_presets():        expected = _label_in_millions(name)        if expected is None:            continue        actual = _total_millions(name)        assert abs(actual - expected) / expected < 0.07, (            f"{name} builds {actual:.1f}M but the name says {expected:.1f}M"        )def test_every_preset_validates():    for name in list_presets():        cfg: ModelConfig = get_preset(name)["model"]        assert cfg.dim % cfg.n_heads == 0        assert cfg.n_heads % cfg.n_kv_heads == 0        assert cfg.rope_theta > 0        assert cfg.max_seq_len > 0def test_dim_must_divide_by_heads():    with pytest.raises(ValueError, match="not divisible"):        ModelConfig(dim=100, n_heads=12, n_kv_heads=4)def test_heads_must_be_multiple_of_kv_heads():    with pytest.raises(ValueError, match="not a multiple"):        ModelConfig(dim=96, n_heads=12, n_kv_heads=5)def test_config_roundtrips_through_dict():    cfg = ModelConfig(dim=64, n_layers=2, n_heads=4, n_kv_heads=2, head_dim=16, ff_dim=128)    assert ModelConfig.from_dict(cfg.to_dict()) == cfgdef test_inject_layers_resolves_negative_indices():    assert AffectConfig(inject_layers=(-1,)).layers_to_inject(12) == (11,)    assert AffectConfig(inject_layers=(-2, -1)).layers_to_inject(12) == (10, 11)    assert AffectConfig(inject_layers=(0, 3)).layers_to_inject(12) == (0, 3)def test_inject_layers_defaults_to_all():    assert AffectConfig().layers_to_inject(6) == tuple(range(6))def test_inject_layers_rejects_out_of_range():    with pytest.raises(ValueError, match="out of range"):        AffectConfig(inject_layers=(99,)).layers_to_inject(12)def test_param_count_helper_agrees_with_module():    preset = get_preset("elafry-tiny")    assert param_count("elafry-tiny") == sum(        p.numel() for p in Elafry(preset["model"], preset["affect"]).parameters()    )def test_8b_is_untied_and_lands_on_8b():    """The 128k-vocab presets drop weight tying, which is the only reason they reach their    stated size. Tying would give 7.50B."""    preset = get_preset("elafry-8b")    assert preset["model"].tie_embeddings is False    total = sum(param_breakdown(preset["model"], preset["affect"]).values())    assert 8.0e9 < total < 8.1e9def test_8b_would_miss_the_mark_if_tied():    """Documents why the untied choice exists, so nobody 'simplifies' it back."""    preset = get_preset("elafry-8b")    tied = ModelConfig.from_dict({**preset["model"].to_dict(), "tie_embeddings": True})    tied_total = sum(param_breakdown(tied, preset["affect"]).values())    assert tied_total < 7.6e9def test_100m_geometry_matches_the_existing_checkpoint():    """Elafry's 100M preset is deliberately the 12-layer/768-dim shape of the checkpoint    already sitting in the parent repo, so the two are directly comparable."""    cfg = get_preset("elafry-100m")["model"]    assert (cfg.n_layers, cfg.dim, cfg.n_heads, cfg.n_kv_heads, cfg.ff_dim) == (        12,        768,        12,        4,        2048,    )def test_state_feature_width_defaults_to_bare_vad():    assert AffectConfig().state_feature_dim == 3    assert AffectConfig(use_intent=True).state_feature_dim == 4    assert AffectConfig(use_intent=True, use_velocity=True, use_crisis=True).state_feature_dim == 6def test_unknown_preset_raises_with_a_helpful_message():    with pytest.raises(KeyError, match="unknown preset"):        get_preset("elafry-69b")