Skip to content

[Fix] Restore backend dispatch under torch.compiler.disable - #1111

Merged
zhiyuan1i merged 2 commits into
fla-org:mainfrom
xy200303:fix-backend-dispatch
Aug 29, 2026
Merged

zhiyuan1i merged 2 commits into
fla-org:mainfrom
xy200303:fix-backend-dispatch

Conversation

@xy200303

@xy200303 xy200303 commented Aug 8, 2026 •

Copy link
Copy Markdown
Contributor

Summary

Fixes #1110. For ops decorated as

@dispatch('kda')
@torch.compiler.disable
def chunk_kda(...):

the dispatch wrapper is silently discarded at import time: functools.wraps copies _torchdynamo_orig_callable from the compiler.disable wrapper onto the dispatch wrapper, so dynamo's innermost_fn unwrapping bypasses dispatch entirely and backend selection never runs. Swapping the decorator order (@torch.compiler.disable outermost) keeps dispatch intact while the call remains excluded from compile graphs.

Affected entry points, all fixed here (one line each):

  • chunk_kda (fla/ops/kda/chunk.py) — the flash_kda backend from [KDA] Support FLASHKDA backend #852 was never actually selected; test_flash_kda_chunk* passed vacuously via the Triton path matching the same fp64 gold
  • fused_kda_gate (fla/ops/kda/gate.py) — no registered backend implements it today, so no runtime change; fixed for consistency
  • chunk_gated_delta_rule (fla/ops/gated_delta_rule/chunk.py) — activates flash_qla dispatch on systems with the package installed (the intent of [GDN] Add FlashQLA backend dispatch #998 / b4cfd53)

Impact / blast radius

This is a behavior fix, not just a refactor: on systems with flash_kda installed, inference-mode chunk_kda calls now genuinely route to the flash_kda backend (it has default_enable=True). That is what #852 intended, but it is an observable change for those users; anyone who needs the old behavior can set FLA_FLASH_KDA=0 (or FLA_DISABLE_BACKEND_DISPATCH=1). Systems without these packages are unaffected — verifiers fall back to the default Triton path, so CI is unchanged.

Test plan

  • tests/ops/test_kda.py — full suite passes; test_flash_kda_chunk* now genuinely exercise the flash_kda backend (verified by spying on the backend entry)
  • tests/ops/test_gdn.py, tests/ops/test_gdn2.py, tests/ops/test_gdn_kernels.py — pass (flash_qla not installed here; dispatch iterates and falls back). One flake: test_gdn2.py::test_chunk_varlen[cu_seqlens[0, 64, 128]-H2-K64-V64-gateFalse-torch.float32] failed once in the full-suite run and passed on isolated rerun with comfortable margins (all ratios ~1e-3); the fix does not alter the GDN code path without flash_qla installed
  • Dependent tests identified via scripts/find_dependent_tests.py for the three touched files

Benchmark / NCU (kernel changes only)

Neutral — decorator order only; no kernel, numerics, or scheduling change. The newly-activated flash_kda inference path is faster than the Triton default (see #852 / FlashKDA benchmarks), not slower.

Checklist

  • I have read CONTRIBUTING.md and follow its conventions (code style, docstrings, commit prefixes).
  • I have read AGENTS.md and, where my change matches its scope, the relevant skill under .agents/skills.
  • This is not a minor/cosmetic-only PR (typo, formatting, style-only tweaks).
  • Dependent tests pass locally or in CI; new behavior is covered by tests where applicable.
  • Kernel changes include same-hardware before/after benchmark numbers (dense + varlen where applicable).

Breaking changes

None for default configurations. Behavior change for users with flash_kda/flash_qla installed, as described above (restores the intended behavior of #852/#998).

xy200303 and others added 2 commits August 8, 2026 14:59
@dispatch stacked over @torch.compiler.disable was silently discarded:
functools.wraps copies _torchdynamo_orig_callable onto the dispatch
wrapper, so dynamo's innermost_fn unwrapping bypasses it and backend
selection never ran for the affected top-level entries:

- chunk_kda (fla/ops/kda/chunk.py) — the flash_kda backend from fla-org#852
  was never actually selected
- fused_kda_gate (fla/ops/kda/gate.py)
- chunk_gated_delta_rule (fla/ops/gated_delta_rule/chunk.py)

Swap the decorator order so dispatch stays the outermost wrapper while
the call remains excluded from compile graphs. Note this activates the
flash_kda backend for inference-mode chunk_kda calls on systems with
the flash_kda package installed, which is the intent of fla-org#852.

Closes fla-org#1110
@zhiyuan1i zhiyuan1i added the bug Something isn't working label Aug 14, 2026
@zhiyuan1i
zhiyuan1i merged commit ad4af37 into fla-org:main Aug 29, 2026
23 of 25 checks passed
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

bug Something isn't working

Projects

None yet

Development

Successfully merging this pull request may close these issues.

[Bug] @dispatch over @torch.compiler.disable is silently discarded — top-level backend dispatch never fires

2 participants