Fix dummy initialization of inverse weight scales (#35491)
Co-authored-by: weireweire <20922698+weireweire@users.noreply.github.com> Co-authored-by: Po-Han Huang (NVIDIA) <53919306+nvpohanh@users.noreply.github.com>
This commit is contained in:
co-authored by
weireweire
Po-Han Huang
parent
e57e934bcc
commit
fb8d7eedda
@@ -1664,8 +1664,11 @@ def initialize_dummy_weights(
|
||||
is fixed, the random values generated by this function only depends on
|
||||
the parameter's number of elements and its data type.
|
||||
"""
|
||||
for param in model.state_dict().values():
|
||||
for name, param in model.state_dict().items():
|
||||
if torch.is_floating_point(param):
|
||||
if name.endswith("weight_scale_inv"):
|
||||
param.fill_(1.0)
|
||||
continue
|
||||
generator = torch.Generator(device=param.data.device)
|
||||
generator.manual_seed(seed)
|
||||
# Tensor subclasses such as MXFP8 wrappers expose a low-bit raw
|
||||
|
||||
Reference in New Issue
Block a user