Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 6 additions & 3 deletions aneforge/llm.py
Original file line number Diff line number Diff line change
Expand Up @@ -269,6 +269,8 @@ def __init__(self, cfg: LlamaConfig, weights: dict, compress: str | None = None,
self.cfg = cfg; self.w = weights; self._net = None; self._seq = 0; self._dec = None; self._pre = None
self.compress = compress # None=fp16, or "int8"/"int4"/"blockwise" to quantize the ANE weights
self._chunk_bytes = 1.6e9 # max fp16 weight bytes per decode program (under the ~2GB ANE ceiling)
self._chunk_max_layers = 24 # max layers per decode program: deep small-dim models stay under the byte
# cap but hit an op-count ceiling (~31 layers, e.g. SmolLM2's 32) -> rc=-1
self._lmT = None # cached contiguous fp32 lm_head^T (host matmul); built on first use
self.ane_lm_head = ane_lm_head # run lm_head on the ANE (tiled compile_multi) instead of host matmul
self._lmh: dict | None = None # cached ANE lm_head program (built once, shape-independent of max_len)
Expand Down Expand Up @@ -421,12 +423,13 @@ def prefill(self, token_ids):
return self._logits(self._hidden(token_ids)[-1])[None]

def _layer_chunks(self):
"""Group layers into contiguous chunks whose baked weights stay under the ANE single-program ceiling
(~2GB; measured ~1.5GB OK, ~3.4GB fails). int8/int4 weights are smaller, so more layers fit per chunk."""
"""Group layers into contiguous chunks under two ANE single-program ceilings: baked weight bytes
(~2GB; measured ~1.5GB OK, ~3.4GB fails; int8/int4 are smaller, so more layers fit) and an op-count
ceiling that a deep small-dim model hits well under the byte cap (`_chunk_max_layers`)."""
per = sum(int(np.asarray(v).size) for v in self.w["layers"][0].values() if isinstance(v, (np.ndarray, list, tuple))) * 2 # fp16 bytes / layer
if self.compress in ("int8", "blockwise"): per //= 2
elif self.compress == "int4": per //= 4
n = max(1, min(self.cfg.n_layers, int(self._chunk_bytes // max(per, 1))))
n = max(1, min(self.cfg.n_layers, self._chunk_max_layers, int(self._chunk_bytes // max(per, 1))))
return [range(i, min(i + n, self.cfg.n_layers)) for i in range(0, self.cfg.n_layers, n)]

def _decoder(self, M):
Expand Down
46 changes: 46 additions & 0 deletions tests/test_decode_chunking.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,46 @@
"""Deep small-dim models: the decode program must chunk by layer count, not just weight bytes.

A model like SmolLM2-360M (32 small layers) stays under the per-program weight ceiling but hits an
op-count ceiling (~31 layers) and the ANE fails the execute with rc=-1. `_layer_chunks` caps layers
per chunk so these segment and decode correctly."""
from aneforge.llm import LlamaConfig
from _helpers import requires_ane, make_random_llama_model as _random_model


def _deep_cfg(n_layers):
return LlamaConfig(dim=128, n_heads=4, n_kv_heads=2, ffn_dim=256, vocab=64, n_layers=n_layers, head_dim=32)


def test_layer_chunks_caps_layer_count():
"""A deep, small-dim model that fits the byte budget must still split on the layer-count cap."""
m = _random_model(_deep_cfg(32)) # 32 small layers: well under _chunk_bytes
chunks = m._layer_chunks()
assert max(len(c) for c in chunks) <= m._chunk_max_layers
assert sum(len(c) for c in chunks) == 32 # every layer is covered exactly once
assert [i for c in chunks for i in c] == list(range(32))


def test_layer_chunks_single_when_shallow():
"""A shallow model stays a single chunk (no needless segmentation / dispatch overhead)."""
m = _random_model(_deep_cfg(8))
assert m._layer_chunks() == [range(0, 8)]


@requires_ane
def test_deep_model_decodes_without_rc_error():
"""32 layers used to fail decode with `execute failed rc=-1`; with the layer cap it generates."""
out = _random_model(_deep_cfg(32)).generate([1, 2, 3], max_new_tokens=3)
assert len(out) == 3


@requires_ane
def test_segmented_decode_matches_single_program():
"""A model split into >1 decode chunk must produce the same greedy tokens as one that fits in a chunk."""
cfg = _deep_cfg(28)
m = _random_model(cfg)
assert len(m._layer_chunks()) > 1 # 28 > cap -> segmented
ref = _random_model(cfg)
ref._chunk_max_layers = 999 # force a single chunk on an identical model
assert len(ref._layer_chunks()) == 1
prompt = [1, 2, 3, 4]
assert list(m.generate(prompt, max_new_tokens=5)) == list(ref.generate(prompt, max_new_tokens=5))