(
files: list[str],
named_tensors: dict[str, torch.Tensor],
rank: int | None = None,
shared_pin_memory: list[MemoryBuffer] | None = None,
)
| 275 | |
| 276 | |
| 277 | def _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, |
no test coverage detected