MCPcopy Create free account
hub / github.com/Pints-AI/1.5-Pints / NotYetLoadedTensor

Class NotYetLoadedTensor

lit_gpt/utils_old.py:94–191  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

92
93
94class 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")

Callers 1

rebuild_tensor_v2Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected