MCPcopy Create free account
hub / github.com/MoonshotAI/checkpoint-engine / _normal_pin_memory

Function _normal_pin_memory

checkpoint_engine/pin_memory.py:277–362  ·  view source on GitHub ↗
(
    files: list[str],
    named_tensors: dict[str, torch.Tensor],
    rank: int | None = None,
    shared_pin_memory: list[MemoryBuffer] | None = None,
)

Source from the content-addressed store, hash-verified

275
276
277def _normal_pin_memory(
278 files: list[str],
279 named_tensors: dict[str, torch.Tensor],
280 rank: int | None = None,
281 shared_pin_memory: list[MemoryBuffer] | None = None,
282) -> list[MemoryBuffer]:
283 parameters = _load_checkpoint(files)
284 if named_tensors:
285 parameters.update(named_tensors)
286 bucket_size = max(4 << 30, max(_align_size(x.dtype, x.shape) for x in parameters.values()))
287
288 class MemoryBucket(BaseModel):
289 size: int
290 metas: list[ParameterMeta]
291
292 buckets: list[MemoryBucket] = []
293 buckets.append(MemoryBucket(size=0, metas=[]))
294 for name, tensor in sorted(parameters.items()):
295 size = _align_size(tensor.dtype, tensor.shape)
296 if buckets[-1].size + size > bucket_size:
297 assert buckets[-1], f"buckets[{len(buckets) - 1}] should not be empty"
298 buckets.append(MemoryBucket(size=0, metas=[]))
299 buckets[-1].metas.append(
300 ParameterMeta(name=name, shape=tensor.shape, dtype=tensor.dtype, aligned_size=size)
301 )
302 buckets[-1].size += size
303
304 memory_buffers = [
305 MemoryBuffer(buffer=torch.empty(0), size=bucket.size, metas=bucket.metas)
306 for bucket in buckets
307 ]
308
309 def register_pin_memory(
310 idx: int, size: int, shared_pin_memory: list[MemoryBuffer] | None = None
311 ) -> tuple[int, torch.Tensor]:
312 if shared_pin_memory:
313 # If shared_pin_memory is provided, reuse the pin memory buffer, do not allocate new one
314 # Reusing pin memory only support fixed shape of checkpoints, which is registered the first time
315 assert idx < len(shared_pin_memory), (
316 f"idx {idx} should be less than shared_pin_memory length {len(shared_pin_memory)}"
317 )
318 assert shared_pin_memory[idx].size == size, (
319 f"shared_pin_memory[{idx}].size {shared_pin_memory[idx].size} should be equal to {size}"
320 )
321 return idx, shared_pin_memory[idx].buffer
322 else:
323 buffer = torch.empty(size, dtype=torch.uint8, pin_memory=True)
324 return idx, buffer
325
326 def register_tensor(buffer: torch.Tensor, offset: int, tensor: torch.Tensor):
327 buffer[offset : offset + tensor.nbytes] = tensor.view(-1).view(dtype=torch.uint8)
328
329 with concurrent.futures.ThreadPoolExecutor(max_workers=32) as executor:
330 futures = [
331 executor.submit(
332 register_pin_memory,
333 idx,
334 bucket.size,

Callers 1

_register_checkpointFunction · 0.85

Calls 6

ParameterMetaClass · 0.90
MemoryBufferClass · 0.90
_load_checkpointFunction · 0.85
_align_sizeFunction · 0.85
MemoryBucketClass · 0.85
updateMethod · 0.80

Tested by

no test coverage detected