Skip to content

LFM2/LFM2.5: short-conv state is not cleared when a new sequence starts, so output after reset() depends on the previous prompt聽#23262

Description

@alpharomercoma

馃悰 Describe the bug

Summary

In exported LFM2 and LFM2.5 models, ShortConv.conv_state (a mutable buffer holding the last two conv inputs) is never cleared when a new sequence starts at input_pos = 0. After TextLLMRunner.reset() (Python), LlmModule.resetContext() (Android) or TextLLMRunner::warmup(), the next conversation starts from the previous conversation's convolution state. Continuing a conversation at the current position is not affected.

  • Expected: after reset(), a prompt gives the same output as on a freshly loaded model.
  • Actual: the output depends on what ran before the reset.

To reproduce

Environment: executorch==1.5.1 and torch==2.14.0 from pip.
Model: a public export, lfm_2_5_350m_xnnpack_8da4w.pte (sha256 e8c38e0d933932b02dddbcf92df813045899d0f71bf5346dbceecf5c4078f05d), and its tokenizer.json (sha256 2221c71b5dce048a8abae62843b92bd7deec13cc153f95fa6e3327a47b79a7da).

1. Module level (python module_repro.py lfm_2_5_350m_xnnpack_8da4w.pte)

import sys
import torch
from executorch.extension.llm.custom_ops import custom_ops  # noqa: F401
from executorch.kernels import quantized  # noqa: F401
from executorch.extension.pybindings.portable_lib import _load_for_executorch

PTE = sys.argv[1]  # e.g. lfm_2_5_350m_xnnpack_8da4w.pte

def run(module, tokens):
    for pos, tok in enumerate(tokens):
        logits = module.forward((torch.tensor([[tok]]), torch.tensor([pos])))[0]
    return logits.reshape(-1).float()

A = [1, 6423, 1098, 4123, 872, 3301, 911]  # an earlier, unrelated sequence
B = [1, 1715, 5902, 3314, 7310]            # a new sequence, positions restart at 0

fresh = run(_load_for_executorch(PTE), B)
reused = _load_for_executorch(PTE)
run(reused, A)
after_a = run(reused, B)

print("max |fresh - after_a| =", (fresh - after_a).abs().max().item())   # expected 0.0
print("argmax fresh / after_a:", fresh.argmax().item(), after_a.argmax().item())

Output:

max |fresh - after_a| = 16.080059051513672
argmax fresh / after_a: 7 3466

Controls: two fresh loads give 0.0, and the same script on an attention-only model (a Qwen3-1.7B XNNPACK export) gives 0.0.

2. Through TextLLMRunner, greedy, with reset() (python runner_repro.py lfm_2_5_350m_xnnpack_8da4w.pte tokenizer.json)

import sys
from executorch.extension.llm.custom_ops import custom_ops  # noqa: F401
from executorch.kernels import quantized  # noqa: F401
from executorch.extension.llm.runner import GenerationConfig, TextLLMRunner

PTE, TOKENIZER = sys.argv[1], sys.argv[2]

def chat(q):
    return f"<|startoftext|><|im_start|>user\n{q}<|im_end|>\n<|im_start|>assistant\n"

def generate(runner, prompt):
    out = []
    runner.generate(prompt, GenerationConfig(echo=False, max_new_tokens=64, temperature=0.0), out.append)
    return "".join(out)

fresh = generate(TextLLMRunner(PTE, TOKENIZER), chat("List five uses for a paperclip."))

runner = TextLLMRunner(PTE, TOKENIZER)
generate(runner, chat("Write a short poem about the ocean at night."))
runner.reset()  # documented as "Reset the runner state and KV cache"
after_reset = generate(runner, chat("List five uses for a paperclip."))

print("identical:", fresh == after_reset)          # expected True
print("fresh      :", repr(fresh[:60]))
print("after reset:", repr(after_reset[:60]))

Last three lines of output (the runner also streams each generation):

identical: False
fresh      : 'Sure! Here are five practical uses for a paperclip:\n\n1. **St'
after reset: 'Here are five uses for a paperclip:\n\n1. **Home Decor**: Pape'

Scope

Each export was checked on 4 random sequence pairs (harness and raw output below). "Leaks" means every pair differed while fresh vs fresh stayed at 0.0.

Export Backend, dtype Result
LFM2-350M, exported with the README's first example XNNPACK 8da4w leaks (4/4)
LFM2.5-350M, same recipe with lfm2_5_350m config (and lfm2_xnnpack_fp32.yaml for fp32) XNNPACK 8da4w and fp32 leaks (4/4 each)
Software Mansion LFM2.5-350M and 1.2B XNNPACK 8da4w leaks (4/4 each)
Software Mansion LFM2.5-1.2B XNNPACK fp16 leaks (4/4)
younghan-meta/LFM2.5-ExecuTorch-MLX 350M MLX 4-bit leaks (4/4)
Qwen3-1.7B (attention only, control) XNNPACK no leak (0/4)

Reproduced on macOS arm64; the Software Mansion 350M file also leaks (4/4) on Linux aarch64 (python:3.12-slim, pip install executorch==1.5.1 torch==2.14.*). A short, confident answer can come out unchanged because the logits shift without always changing the argmax; open-ended prompts diverge.

Root cause

Permalinks at 2f78245. short_conv.py is identical in v1.4.0 and v1.5.1, and so are the TextLLMRunner::reset() and warmup() bodies.

  • conv_state is registered at short_conv.py#L42, then prepended and overwritten on every call at #L66-L75. The comment at #L63 assumes prefill starts on an empty cache.
  • ShortConvBlock.forward receives the attention options that carry input_pos but ignores them (#L104). reset_cache() (#L85) is never called outside eager code.
  • TextLLMRunner::reset() (text_llm_runner.cpp#L354) resets stats, pos_ and the pending prefill token only, and warmup() (#L331) calls it after generating. The Python binding documents reset() as "Reset the runner state and KV cache" (pybindings.cpp#L735).

Verified fix

Qwen3.5's GatedDeltaNet already clears its state when input_pos[0] == 0 (_maybe_reset_state, attention.py#L801-L811). Applying the same pattern to ShortConv and re-exporting with the same recipe removes the leak:

def conv_forward(self, x, input_pos=None):              # ShortConv
    if input_pos is not None:
        self.conv_state.mul_(1.0 - (input_pos[0] == 0).to(self.conv_state.dtype))
    return original_conv_forward(self, x)

def block_forward(self, x, freqs_cos=None, freqs_sin=None, _unused_attn_options=None):   # ShortConvBlock
    input_pos = (_unused_attn_options or {}).get("input_pos")
    h = x + self.conv.forward(self.attention_norm(x), input_pos)
    return h + self.feed_forward(self.ffn_norm(h)), None
Check Unpatched Patched
Repro 1, max logit difference 15.55 0.0
Repro 2 identical: False identical: True
Random pairs leaking 4/4 0/4
Fresh sequence, patched vs unpatched logits identical (0.0)

(LFM2.5-350M fp32 export, same recipe. The patch is a no-op while input_pos[0] != 0 and clears conv_state whenever a sequence starts at position 0, with or without reset().)

Suggested next steps

  1. Apply the patch above and add a regression test like Qwen3.5's test_gated_deltanet_resets_state_on_new_sequence. If input_pos=None is supported, reset state on that path too, as _maybe_reset_state does.
  2. Files already exported keep the leak until re-exported, so clearing mutable buffers in reset() may also be worth considering.

I can open a PR with the model-side change and a test.

Related

Harness and raw output (all exports)
"""Quadruple check of the LFM2 state leak. Module-level: does a new sequence at input_pos 0 on a reused module give
the same logits as the same sequence on a freshly loaded module? Several random sequence pairs, several lengths.

  python harness.py <file.pte> [vocab_cap]
"""
import random, sys, hashlib
import torch
from executorch.extension.llm.custom_ops import custom_ops  # noqa: F401
from executorch.kernels import quantized  # noqa: F401
from executorch.extension.pybindings.portable_lib import _load_for_executorch

pte = sys.argv[1]
cap = int(sys.argv[2]) if len(sys.argv) > 2 else 8000
sha = hashlib.sha256(open(pte, "rb").read()).hexdigest()[:16]

def run(m, toks):
    out = None
    for i, t in enumerate(toks):
        out = m.forward((torch.tensor([[t]], dtype=torch.long), torch.tensor([i], dtype=torch.long)))[0]
    return out.reshape(-1).float()

def kl(a, b):
    p, q = torch.log_softmax(a, -1), torch.log_softmax(b, -1)
    return float((p.exp() * (p - q)).sum())

m = _load_for_executorch(pte)
print(f"{pte.split('/')[-1]} sha256:{sha} methods={[x for x in m.method_names() if x == 'forward' or 'reset' in x]}")
rng = random.Random(0)
rows = []
for trial, (la, lb) in enumerate([(1, 4), (6, 5), (12, 8), (3, 3)]):
    A = [1] + [rng.randrange(100, cap) for _ in range(la)]
    B = [1] + [rng.randrange(100, cap) for _ in range(lb)]
    fresh = run(_load_for_executorch(pte), B)
    fresh2 = run(_load_for_executorch(pte), B)
    m2 = _load_for_executorch(pte); run(m2, A); after = run(m2, B)
    m3 = _load_for_executorch(pte); run(m3, B); again = run(m3, B)     # same sequence twice on one module
    rows.append((trial, la, lb, float((fresh - fresh2).abs().max()), float((fresh - after).abs().max()),
                 int(fresh.argmax()) != int(after.argmax()), kl(fresh, after), float((fresh - again).abs().max())))
for r in rows:
    print("  trial %d |A|=%-2d |B|=%-2d  control %.6f | A-then-B max %.4f argmax-changed %-5s KL %.4f | B-then-B max %.4f" % r)
leak = sum(r[4] > 1e-3 for r in rows)
print(f"  => {leak}/{len(rows)} trials leak; control max {max(r[3] for r in rows):.6f}")
# harness.py: 4 random sequence pairs per file; executorch 1.5.1 (pip), torch 2.14.0, macOS 26.5.2 arm64
swm_350m_8da4w.pte sha256:e8c38e0d933932b0 methods=['forward']
  trial 0 |A|=1  |B|=4   control 0.000000 | A-then-B max 16.8667 argmax-changed True  KL 7.7509 | B-then-B max 19.7846
  trial 1 |A|=6  |B|=5   control 0.000000 | A-then-B max 18.1271 argmax-changed True  KL 5.0337 | B-then-B max 14.6948
  trial 2 |A|=12 |B|=8   control 0.000000 | A-then-B max 15.9862 argmax-changed True  KL 5.9448 | B-then-B max 17.6679
  trial 3 |A|=3  |B|=3   control 0.000000 | A-then-B max 14.2549 argmax-changed True  KL 2.6568 | B-then-B max 15.3156
  => 4/4 trials leak; control max 0.000000
swm_1_2b_8da4w.pte sha256:69b9d1d7f7d576e1 methods=['forward']
  trial 0 |A|=1  |B|=4   control 0.000000 | A-then-B max 6.5088 argmax-changed True  KL 1.1331 | B-then-B max 7.9537
  trial 1 |A|=6  |B|=5   control 0.000000 | A-then-B max 10.9939 argmax-changed True  KL 2.5087 | B-then-B max 9.8512
  trial 2 |A|=12 |B|=8   control 0.000000 | A-then-B max 4.0695 argmax-changed False KL 0.2427 | B-then-B max 8.2273
  trial 3 |A|=3  |B|=3   control 0.000000 | A-then-B max 6.7117 argmax-changed True  KL 1.8102 | B-then-B max 9.4184
  => 4/4 trials leak; control max 0.000000
swm_1_2b_fp16.pte sha256:cd013f51d5f7f3a7 methods=['forward']
  trial 0 |A|=1  |B|=4   control 0.000000 | A-then-B max 8.8141 argmax-changed True  KL 1.5633 | B-then-B max 9.9648
  trial 1 |A|=6  |B|=5   control 0.000000 | A-then-B max 9.6172 argmax-changed True  KL 1.7002 | B-then-B max 7.9590
  trial 2 |A|=12 |B|=8   control 0.000000 | A-then-B max 3.7124 argmax-changed False KL 0.2932 | B-then-B max 10.9331
  trial 3 |A|=3  |B|=3   control 0.000000 | A-then-B max 5.7681 argmax-changed False KL 0.5080 | B-then-B max 6.1289
  => 4/4 trials leak; control max 0.000000
yh_350m_mlx_4w.pte sha256:a22641ff8364b813 methods=['forward']
  trial 0 |A|=1  |B|=4   control 0.000000 | A-then-B max 19.9961 argmax-changed True  KL 5.9034 | B-then-B max 18.8750
  trial 1 |A|=6  |B|=5   control 0.000000 | A-then-B max 18.4062 argmax-changed True  KL 5.4028 | B-then-B max 15.4062
  trial 2 |A|=12 |B|=8   control 0.000000 | A-then-B max 16.9062 argmax-changed True  KL 8.4498 | B-then-B max 15.3125
  trial 3 |A|=3  |B|=3   control 0.000000 | A-then-B max 16.5312 argmax-changed True  KL 3.9572 | B-then-B max 15.4570
  => 4/4 trials leak; control max 0.000000
fresh_lfm2_5_350m_8da4w.pte sha256:e7e1fa1c7f357243 methods=['forward']
  trial 0 |A|=1  |B|=4   control 0.000000 | A-then-B max 17.0455 argmax-changed True  KL 7.5597 | B-then-B max 18.3524
  trial 1 |A|=6  |B|=5   control 0.000000 | A-then-B max 14.5359 argmax-changed True  KL 3.8377 | B-then-B max 13.2261
  trial 2 |A|=12 |B|=8   control 0.000000 | A-then-B max 16.6925 argmax-changed True  KL 6.4640 | B-then-B max 13.8776
  trial 3 |A|=3  |B|=3   control 0.000000 | A-then-B max 15.6679 argmax-changed True  KL 3.1766 | B-then-B max 19.1684
  => 4/4 trials leak; control max 0.000000
fresh_lfm2_5_350m_fp32.pte sha256:ed48c2bacfc9c040 methods=['forward']
  trial 0 |A|=1  |B|=4   control 0.000000 | A-then-B max 17.0677 argmax-changed True  KL 5.6892 | B-then-B max 21.2208
  trial 1 |A|=6  |B|=5   control 0.000000 | A-then-B max 19.8826 argmax-changed False KL 4.5853 | B-then-B max 16.9365
  trial 2 |A|=12 |B|=8   control 0.000000 | A-then-B max 19.6481 argmax-changed True  KL 8.9908 | B-then-B max 14.9544
  trial 3 |A|=3  |B|=3   control 0.000000 | A-then-B max 11.9439 argmax-changed True  KL 2.5607 | B-then-B max 17.1218
  => 4/4 trials leak; control max 0.000000
patched_lfm2_5_350m_fp32.pte sha256:d46fe088f609b2ba methods=['forward']
  trial 0 |A|=1  |B|=4   control 0.000000 | A-then-B max 0.0000 argmax-changed False KL 0.0000 | B-then-B max 0.0000
  trial 1 |A|=6  |B|=5   control 0.000000 | A-then-B max 0.0000 argmax-changed False KL 0.0000 | B-then-B max 0.0000
  trial 2 |A|=12 |B|=8   control 0.000000 | A-then-B max 0.0000 argmax-changed False KL 0.0000 | B-then-B max 0.0000
  trial 3 |A|=3  |B|=3   control 0.000000 | A-then-B max 0.0000 argmax-changed False KL 0.0000 | B-then-B max 0.0000
  => 0/4 trials leak; control max 0.000000
model.pte sha256:074d87c3cba37ed8 methods=['forward']
  trial 0 |A|=1  |B|=4   control 0.000000 | A-then-B max 0.0000 argmax-changed False KL 0.0000 | B-then-B max 0.0000
  trial 1 |A|=6  |B|=5   control 0.000000 | A-then-B max 0.0000 argmax-changed False KL 0.0000 | B-then-B max 0.0000
  trial 2 |A|=12 |B|=8   control 0.000000 | A-then-B max 0.0000 argmax-changed False KL 0.0000 | B-then-B max 0.0000
  trial 3 |A|=3  |B|=3   control 0.000000 | A-then-B max 0.0000 argmax-changed False KL 0.0000 | B-then-B max 0.0000
  => 0/4 trials leak; control max 0.000000

# Linux aarch64 (python:3.12-slim container, pip install executorch==1.5.1 torch==2.14.*)
Linux-6.8.0-117-generic-aarch64-with-glibc2.41 torch 2.14.0+cu130
swm_350m_8da4w.pte sha256:e8c38e0d933932b0 methods=['forward']
  trial 0 |A|=1  |B|=4   control 0.000000 | A-then-B max 16.8667 argmax-changed True  KL 7.7509 | B-then-B max 19.7846
  trial 1 |A|=6  |B|=5   control 0.000000 | A-then-B max 18.1271 argmax-changed True  KL 5.0337 | B-then-B max 14.6948
  trial 2 |A|=12 |B|=8   control 0.000000 | A-then-B max 16.0775 argmax-changed True  KL 5.9302 | B-then-B max 17.6679
  trial 3 |A|=3  |B|=3   control 0.000000 | A-then-B max 14.2549 argmax-changed True  KL 2.6568 | B-then-B max 15.3156
  => 4/4 trials leak; control max 0.000000

# LFM2-350M exported with the README's first example command verbatim (lfm2_xnnpack_q8da4w.yaml, lfm2_350m)
lfm2_350m_8da4w.pte sha256:adf2915039c60deb methods=['forward']
  trial 0 |A|=1  |B|=4   control 0.000000 | A-then-B max 9.1079 argmax-changed True  KL 2.4930 | B-then-B max 7.7957
  trial 1 |A|=6  |B|=5   control 0.000000 | A-then-B max 11.9951 argmax-changed True  KL 3.9635 | B-then-B max 12.0744
  trial 2 |A|=12 |B|=8   control 0.000000 | A-then-B max 12.3841 argmax-changed True  KL 3.1355 | B-then-B max 12.4198
  trial 3 |A|=3  |B|=3   control 0.000000 | A-then-B max 11.2085 argmax-changed True  KL 3.0994 | B-then-B max 10.8808
  => 4/4 trials leak; control max 0.000000

Versions

executorch 1.5.1 (pip), torch 2.14.0, torchao 0.18.0, Python 3.12.13, macOS 26.5.2 arm64. The Software Mansion 350M result also reproduces on Linux aarch64 (python:3.12-slim container, pip install executorch==1.5.1 torch==2.14.*). Source checked at main 2f7824593d0d14f9d5d73faa540b681b6b2b507c, v1.5.1 and v1.4.0.

collect_env
Collecting environment information...
PyTorch version: 2.14.0
Is debug build: False
CUDA used to build PyTorch: None
ROCM used to build PyTorch: N/A

OS: macOS 26.5.2 (arm64)
GCC version: Could not collect
Clang version: 21.0.0 (clang-2100.1.1.101)
CMake version: version 4.4.0
Libc version: N/A

Python version: 3.12.13 (main, Jun 23 2026, 15:44:24) [Clang 22.1.3 ] (64-bit runtime)
Python platform: macOS-26.5.2-arm64-arm-64bit
Is CUDA available: False
CUDA runtime version: No CUDA
CUDA_MODULE_LOADING set to: N/A
GPU models and configuration: No CUDA
Nvidia driver version: No CUDA
cuDNN version: No CUDA
Is XPU available: False
HIP runtime version: N/A
MIOpen runtime version: N/A
Is XNNPACK available: False
Caching allocator config: N/A

CPU:
Apple M5

Versions of relevant libraries:
[pip3] Could not collect
[conda] Could not collect
executorch             1.5.1
numpy                  2.5.3
pytorch-tokenizers     1.5.0
torch                  2.14.0
torchao                0.18.0

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions