| 92 | |
| 93 | |
| 94 | class NotYetLoadedTensor: |
| 95 | def __init__(self, metatensor, archiveinfo, storageinfo, rebuild_args): |
| 96 | self.metatensor = metatensor |
| 97 | self.archiveinfo = archiveinfo |
| 98 | self.storageinfo = storageinfo |
| 99 | self.rebuild_args = rebuild_args |
| 100 | |
| 101 | @classmethod |
| 102 | def rebuild_from_type_v2(cls, func, new_type, args, state, *, archiveinfo=None): |
| 103 | ret = func(*args) |
| 104 | if isinstance(ret, NotYetLoadedTensor): |
| 105 | old_lt = ret._load_tensor |
| 106 | |
| 107 | def _load_tensor(): |
| 108 | t = old_lt() |
| 109 | return torch._tensor._rebuild_from_type_v2(lambda: t, new_type, (), state) |
| 110 | |
| 111 | ret._load_tensor = _load_tensor |
| 112 | return ret |
| 113 | return torch._tensor._rebuild_from_type_v2(func, new_type, args, state) |
| 114 | |
| 115 | @classmethod |
| 116 | def rebuild_parameter(cls, data, requires_grad, backward_hooks, *, archiveinfo=None): |
| 117 | if isinstance(data, NotYetLoadedTensor): |
| 118 | old_lt = data._load_tensor |
| 119 | |
| 120 | def _load_tensor(): |
| 121 | t = old_lt() |
| 122 | return torch._utils._rebuild_parameter(t, requires_grad, backward_hooks) |
| 123 | |
| 124 | data._load_tensor = _load_tensor |
| 125 | return data |
| 126 | return torch._utils._rebuild_parameter(data, requires_grad, backward_hooks) |
| 127 | |
| 128 | @classmethod |
| 129 | def rebuild_tensor_v2( |
| 130 | cls, storage, storage_offset, size, stride, requires_grad, backward_hooks, metadata=None, *, archiveinfo=None |
| 131 | ): |
| 132 | rebuild_args = (storage_offset, size, stride, requires_grad, backward_hooks, metadata) |
| 133 | metatensor = torch._utils._rebuild_tensor_v2( |
| 134 | storage, storage_offset, size, stride, requires_grad, backward_hooks, metadata |
| 135 | ) |
| 136 | storageinfo = storage.archiveinfo |
| 137 | return NotYetLoadedTensor(metatensor, archiveinfo, storageinfo, rebuild_args) |
| 138 | |
| 139 | def _load_tensor(self): |
| 140 | name, storage_cls, fn, device, size = self.storageinfo |
| 141 | dtype = self.metatensor.dtype |
| 142 | |
| 143 | uts = ( |
| 144 | self.archiveinfo.zipfile_context.zf.get_storage_from_record( |
| 145 | f"data/{fn}", size * torch._utils._element_size(dtype), torch.UntypedStorage |
| 146 | ) |
| 147 | ._typed_storage() |
| 148 | ._untyped_storage |
| 149 | ) |
| 150 | with warnings.catch_warnings(): |
| 151 | warnings.simplefilter("ignore") |