MCPcopy Create free account
hub / github.com/Francis-Rings/FlashPortrait / replace_parameters_by_name

Function replace_parameters_by_name

wan/utils/fp8_optimization.py:8–17  ·  view source on GitHub ↗
(module, name_keywords, device)

Source from the content-addressed store, hash-verified

6import torch.nn as nn
7
8def replace_parameters_by_name(module, name_keywords, device):
9 from torch import nn
10 for name, param in list(module.named_parameters(recurse=False)):
11 if any(keyword in name for keyword in name_keywords):
12 if isinstance(param, nn.Parameter):
13 tensor = param.data
14 delattr(module, name)
15 setattr(module, name, tensor.to(device=device))
16 for child_name, child_module in module.named_children():
17 replace_parameters_by_name(child_module, name_keywords, device)
18
19def convert_model_weight_to_float8(model, exclude_module_name=['embed_tokens'], device=None):
20 for name, module in model.named_modules():

Callers 2

fast_infer.pyFile · 0.90
infer.pyFile · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected