[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:
co-authored by
Claude Opus 4.6
parent
e6071e60c0
commit
62a63eeff7
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user