"""Exact saved-tensor placement with a bounded GPU retention budget.""" import torch def storage_key(tensor): return (str(tensor.device), tensor.untyped_storage().data_ptr()) class BudgetedSavedTensors(torch.autograd.graph.saved_tensors_hooks): """Keep existing backbone storage and a bounded set of additional storages. Count whole unique storages, not views. Stored tensors are detached aliases; no casting, quantization or recomputation changes backward values. """ def __init__(self, resident_tensors=(), keep_bytes=0): if keep_bytes < 0: raise ValueError('saved-tensor GPU budget must be nonnegative') resident = {storage_key(t) for t in resident_tensors} kept = set() self.retained_bytes = 0 self.resident_hits = 0 self.cpu_copies = 0 cpu = torch.autograd.graph.save_on_cpu(pin_memory=True, device_type='cuda') def pack(tensor): key = storage_key(tensor) size = tensor.untyped_storage().nbytes() if key in resident: self.resident_hits += 1 return ('resident', tensor.detach()) if key in kept or self.retained_bytes + size <= keep_bytes: if key not in kept: kept.add(key) self.retained_bytes += size return ('resident', tensor.detach()) self.cpu_copies += 1 return ('cpu', cpu.pack_hook(tensor)) def unpack(packed): kind, value = packed return value if kind == 'resident' else cpu.unpack_hook(value) super().__init__(pack, unpack) class AsyncSavedTensors(torch.autograd.graph.saved_tensors_hooks): """Pinned D2H without a host synchronization at every saved tensor. An event explicitly orders each H2D restore after its D2H copy, including when autograd executes that restore on a different CUDA stream. """ def __init__(self): def pack(tensor): if tensor.device.type != 'cuda': return (tensor.device, tensor.detach(), None) packed = torch.empty(tensor.size(), dtype=tensor.dtype, layout=tensor.layout, pin_memory=True) stream = torch.cuda.current_stream(tensor.device) packed.copy_(tensor, non_blocking=True) tensor.record_stream(stream) ready = torch.cuda.Event() ready.record(stream) return (tensor.device, packed, ready) def unpack(saved): device, tensor, ready = saved if ready is None: return tensor torch.cuda.current_stream(device).wait_event(ready) return tensor.to(device, non_blocking=True) super().__init__(pack, unpack)