[Fix] Fix weight_loader property assignment for qwen3-next FP8 models (#21662)

Co-authored-by: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
Baizhou Zhang
2026-03-30 01:35:59 -07:00
committed by GitHub
co-authored by Claude Opus 4.6
parent e6071e60c0
commit 62a63eeff7
+17 -4
View File
@@ -135,11 +135,11 @@ class Qwen3GatedDeltaNet(nn.Module):
# Override weight_loader for packed checkpoint format.
# Must capture original_loader BEFORE overwriting.
self.in_proj_qkvz.weight.weight_loader = self._make_packed_weight_loader(
self.in_proj_qkvz
self._override_weight_loader(
self.in_proj_qkvz, self._make_packed_weight_loader(self.in_proj_qkvz)
)
self.in_proj_ba.weight.weight_loader = self._make_packed_weight_loader(
self.in_proj_ba
self._override_weight_loader(
self.in_proj_ba, self._make_packed_weight_loader(self.in_proj_ba)
)
# Conv1d weight loader setup
@@ -216,6 +216,19 @@ class Qwen3GatedDeltaNet(nn.Module):
dt_bias=self.dt_bias,
)
@staticmethod
def _override_weight_loader(module, new_loader):
"""Override weight_loader on a module's weight parameter.
ModelWeightParameter exposes weight_loader as a read-only property
backed by _weight_loader, while plain parameters store it as a
regular attribute. This helper handles both cases."""
param = module.weight
if hasattr(param, "_weight_loader"):
param._weight_loader = new_loader
else:
param.weight_loader = new_loader
@staticmethod
def _make_packed_weight_loader(module):
"""Create a weight_loader that does contiguous TP slicing for fused