MCPcopy Create free account
hub / github.com/OpenImagingLab/FlashVSR / init_weights_on_device

Function init_weights_on_device

diffsynth/models/utils.py:7–53  ·  view source on GitHub ↗
(device = torch.device("meta"), include_buffers :bool = False)

Source from the content-addressed store, hash-verified

5
6@contextmanager
7def init_weights_on_device(device = torch.device("meta"), include_buffers :bool = False):
8
9 old_register_parameter = torch.nn.Module.register_parameter
10 if include_buffers:
11 old_register_buffer = torch.nn.Module.register_buffer
12
13 def register_empty_parameter(module, name, param):
14 old_register_parameter(module, name, param)
15 if param is not None:
16 param_cls = type(module._parameters[name])
17 kwargs = module._parameters[name].__dict__
18 kwargs["requires_grad"] = param.requires_grad
19 module._parameters[name] = param_cls(module._parameters[name].to(device), **kwargs)
20
21 def register_empty_buffer(module, name, buffer, persistent=True):
22 old_register_buffer(module, name, buffer, persistent=persistent)
23 if buffer is not None:
24 module._buffers[name] = module._buffers[name].to(device)
25
26 def patch_tensor_constructor(fn):
27 def wrapper(*args, **kwargs):
28 kwargs["device"] = device
29 return fn(*args, **kwargs)
30
31 return wrapper
32
33 if include_buffers:
34 tensor_constructors_to_patch = {
35 torch_function_name: getattr(torch, torch_function_name)
36 for torch_function_name in ["empty", "zeros", "ones", "full"]
37 }
38 else:
39 tensor_constructors_to_patch = {}
40
41 try:
42 torch.nn.Module.register_parameter = register_empty_parameter
43 if include_buffers:
44 torch.nn.Module.register_buffer = register_empty_buffer
45 for torch_function_name in tensor_constructors_to_patch.keys():
46 setattr(torch, torch_function_name, patch_tensor_constructor(getattr(torch, torch_function_name)))
47 yield
48 finally:
49 torch.nn.Module.register_parameter = old_register_parameter
50 if include_buffers:
51 torch.nn.Module.register_buffer = old_register_buffer
52 for torch_function_name, old_torch_function in tensor_constructors_to_patch.items():
53 setattr(torch, torch_function_name, old_torch_function)
54
55def load_state_dict_from_folder(file_path, torch_dtype=None):
56 state_dict = {}

Callers 5

__init__Method · 0.85
replace_layerMethod · 0.85
replace_layerMethod · 0.85
replace_layerMethod · 0.85

Calls 1

patch_tensor_constructorFunction · 0.85

Tested by

no test coverage detected