linkedin / linkedin/Liger-Kernel
`torch.compile` fails with FLCE due to `aten.addmm.dtype_out` overload
Nobody has claimed this yet.
- Dominant language
- Python
- Stars
- 6.6k
- Forks
- 603
- Avg merge
- 1d 20h
- Merged PRs (30d)
- 47
Description
## Description
`LigerFusedLinearCrossEntropy` (FLCE) raises an error when traced via `torch.compile` under the following conditions:
- PyTorch 2.8.0 or later.
- NVIDIA CUDA, compute capability 8.0 or higher.
- The gradient computation path for `grad_weight` passes through `torch.addmm(..., out_dtype=..., out=...)`.
- PyTorch uses FakeTensor to trace the graph.
The error occurs at trace time, before any CUDA kernel is executed. Eager mode is unaffected.
## Root Cause
When both `out_dtype=` and `out=` are passed simultaneously, the dispatcher selects the overload:
```
aten.addmm.dtype_out
```
This is a two-tier issue with different fix statuses at each tier:
### Tier 1 — Core decomposition (Fixed on `main` & PyTorch 2.13.0)
The legacy decomposition registered for the entire `aten.addmm` packet used the `@pw_cast_for_opmath` decorator, which remaps positional arguments after upcasting. Consequently, `out_dtype` (of type `torch.dtype`) was passed into the position of `beta`, making the internal expression `beta * self` evaluate as `torch.dtype * FakeTensor` and raising:
```
TypeError: unsupported operand type(s) for *: 'torch.dtype' and 'FakeTensor'
```
[PR #179634](https://github.com/pytorch/pytorch/pull/179634) (commit [`c4c886d`](https://github.com/pytorch/pytorch/commit/c4c886d9f62b69150800e0fc96199e3b36d61181)) created a dedicated `addmm_dtype` handler for overloads containing `out_dtype` without using `@pw_cast_for_opmath`. This commit is an ancestor of `main` and is included in PyTorch 2.13.0.
### Tier 2 — TorchInductor (Unfixed)
TorchInductor maintains its own decomposition table and currently registers decompositions across the full `aten.addmm` packet, preventing the fixed Tier 1 core decomposition from handling `.dtype` and `.dtype_out` overloads. While the operation still fails under `torch.compile`, the failure signature following the core fix may change—for instance, shifting to a lowering error such as:
```
TypeError: tuned_addmm() takes 3 positional arguments but 4 were given
```
The complete fix for Tier 2 is currently tracked in [PR #190936](https://github.com/pytorch/pytorch/pull/190936) (open, unmerged, CI failing). The original bug was reported in [pytorch/pytorch#163880](https://github.com/pytorch/pytorch/issues/163880) (currently marked *Closed*, but outdated—the issue closed automatically following a partial fix that was subsequently reverted, while the bug remains reproducible on `main`).
## Scope of Impact on `main`
The affected fast path is triggered when all of the following conditions are met:
- PyTorch ≥ 2.8 (including 2.13.0).
- NVIDIA CUDA, compute capability ≥ 8.0.
- `grad_weight` has FP32 dtype.
- `grad_logits` has FP16 or BF16 dtype.
- Weight gradients are required (does not apply to inference-only, frozen weights, or passes where weight gradients are skipped).
- Code is executed via `torch.compile` / TorchInductor.
Configuring `accum_dtype=torch.float32` with FP16/BF16 inputs is the minimal setup to reproduce this issue.
## How to Reproduce
Run the following script in an environment with Liger-Kernel installed:
```python
import contextlib
import io
import logging
import platform
from importlib.metadata import PackageNotFoundError, version
import torch
KNOWN_FAILURES = (
# Core decomposition error — PyTorch 2.8
"unsupported operand type(s) for *: 'torch.dtype' and 'FakeTensor'",
# Inductor lowering error — PyTorch 2.13+ / main after core fix
"tuned_addmm() takes 3 positional arguments but 4 were given",
)
def package_version(name, fallback):
try:
return version(name)
except PackageNotFoundError:
return fallback
def main():
cuda_available = torch.cuda.is_available()
print(f"OS: {platform.platform()}")
print(f"Python: {platform.python_version()}")
print(f"PyTorch: {torch.__version__}")
print(f"Liger: {package_version('liger-kernel', 'source checkout')}")
print(f"Triton: {package_version('triton', 'not installed')}")
print(f"CUDA: {torch.version.cuda or 'not built'}")
if cuda_available:
capability = torch.cuda.get_device_capability()
print(
f"GPU: {torch.cuda.get_device_name()} "
f"(compute capability {capability[0]}.{capability[1]})"
)
print("Case: FLCE BF16 operands -> FP32 accumulator, torch.compile/Inductor")
if "dtype_out" not in torch.ops.aten.addmm.overloads():
print("Result: SKIP — requires PyTorch 2.8+")
return 2
if not cuda_available or torch.version.hip is not None:
print("Result: SKIP — requires NVIDIA CUDA")
return 2
if capability[0] < 8:
print("Result: SKIP — requires compute capability 8.0+")
return 2
from liger_kernel.transformers.functional import liger_fused_linear_cross_entropy
x = torch.randn(8, 16, device="cuda", dtype=torch.bfloat16, requires_grad=True)
weight = torch.randn(32, 16, device="cuda", dtype=torch.bfloat16, requires_grad=True)
target = torch.randint(32, (8,), device="cuda")
def step(x, weight, target):
return liger_fused_linear_cross_entropy(
x,
weight,
target,
accum_dtype=torch.float32,
)
logging.disable(logging.CRITICAL)
torch._dynamo.config.capture_scalar_outputs = True
try:
with contextlib.redirect_stderr(io.StringIO()):
torch.compile(step)(x, weight, target).backward()
except Exception as error:
message = str(error)
for expected in KNOWN_FAILURES:
if expected in message:
print(f"Result: REPRODUCED — {message.splitlines()[0]}")
return 0
print(
f"Result: UNEXPECTED — "
f"{type(error).__name__}: {message.splitlines()[0]}"
)
return 1
print("Result: NOT REPRODUCED")
return 1
raise SystemExit(main())
```
## Actual Results on PyTorch 2.8.0
Reproduced on two separate GPUs. The observed failure mode is the Tier 1 core decomposition error characteristic of PyTorch 2.8.0 prior to the core fix merge.
### NVIDIA L4
```
OS: Linux-4.19.0-gvisor-x86_64-with-glibc2.36
Python: 3.12.6
PyTorch: 2.8.0+cu129
Liger: 0.8.1
Triton: 3.4.0
CUDA: 12.9
GPU: NVIDIA L4 (compute capability 8.9)
Case: FLCE BF16 operands -> FP32 accumulator, torch.compile/Inductor
Result: REPRODUCED — unsupported operand type(s) for *: 'torch.dtype' and 'FakeTensor'
```
### NVIDIA A100
```
OS: Linux-4.19.0-gvisor-x86_64-with-glibc2.36
Python: 3.12.6
PyTorch: 2.8.0+cu129
Liger: 0.8.1
Triton: 3.4.0
CUDA: 12.9
GPU: NVIDIA A100-SXM4-40GB (compute capability 8.0)
Case: FLCE BF16 operands -> FP32 accumulator, torch.compile/Inductor
Result: REPRODUCED — unsupported operand type(s) for *: 'torch.dtype' and 'FakeTensor'
```
## Expanded Scope from PR #1324
PR [linkedin/Liger-Kernel#1324](https://github.com/linkedin/Liger-Kernel/pull/1324) (*[Perf] Use direct FLCE weight-gradient accumulation for FP16, BF16, and FP32*, open, unmerged) extends the fast path to same-dtype cases (FP16/FP16, BF16/BF16, FP32/FP32):
```python
torch.addmm(
grad_weight,
grad_logits_t,
input_chunk,
out_dtype=grad_weight.dtype,
out=grad_weight,
)
```
This invocation still dispatches to `aten.addmm.dtype_out` and triggers the same bug. If PR #1324 is merged, the default configuration (`accum_dtype=None`) will enter the `aten.addmm.dtype_out` overload for eligible same-dtype FP16, BF16, or FP32 training workloads. As a result, many callers using `torch.compile` will encounter errors without needing to opt in via `accum_dtype=torch.float32`, whereas the issue is currently limited to that opt-in setup.
Contributor guide
First steps
- Read the whole issue, then the project's contributing guide.
- Comment on the issue to say you are picking it up — it saves two people doing the same work.
- Fork the repository and make your change on a branch.
- Open a pull request that references the issue number.
Research direction
Start with the supplied reproduction script and the liger_kernel.transformers.functional.liger_fused_linear_cross_entropy entry point, then compare the behavior across the listed PyTorch versions. Check the interaction with the PyTorch and TorchInductor fixes referenced in PRs #179634 and #190936; done means the compiled FLCE backward path no longer produces either known failure.
Written by the indexing model from the issue text.
Assessment
- Tech stack
- python, pytorch
- Domain
- machine-learning, performance
- Issue type
- Bug
- Difficulty
- 4/5
- Estimated time
- 3-5 days
- Activity status
- Quiet
- Clarity
- Mostly clear
- Newbie friendliness
- 45/100