(module, state_dict, prefix)
| 145 | return module.__class__ in load_layers or module._get_name() in load_layer_names |
| 146 | |
| 147 | def load_buffer(module, state_dict, prefix): |
| 148 | for name in module._buffers.keys(): |
| 149 | if module._buffers[name].data.is_meta: |
| 150 | module._buffers[name] = torch.nn.parameter.Parameter( |
| 151 | data=torch.empty_like(module._buffers[name].data, device="cpu"), |
| 152 | requires_grad=module._buffers[name].data.requires_grad) |
| 153 | if prefix + name in state_dict.keys(): |
| 154 | module._buffers[name].data.copy_(state_dict[prefix + name]) |
| 155 | |
| 156 | def load(module, state_dict, prefix, mp_group=None): |
| 157 | mp_replace = ReplaceWithTensorSlicing(mp_group=mp_group) |
no test coverage detected