(module, name, param)
| 9 | old_register_buffer = torch.nn.Module.register_buffer |
| 10 | |
| 11 | def register_empty_parameter(module, name, param): |
| 12 | old_register_parameter(module, name, param) |
| 13 | if param is not None: |
| 14 | param_cls = type(module._parameters[name]) |
| 15 | kwargs = module._parameters[name].__dict__ |
| 16 | kwargs["requires_grad"] = param.requires_grad |
| 17 | module._parameters[name] = param_cls( |
| 18 | module._parameters[name].to(device), **kwargs |
| 19 | ) |
| 20 | |
| 21 | def register_empty_buffer(module, name, buffer, persistent=True): |
| 22 | old_register_buffer(module, name, buffer, persistent=persistent) |
nothing calls this directly
no outgoing calls
no test coverage detected