import ctypes import gc import torch from .manager_modules import ( LinearLayerMemoryManager, ConvLayerMemoryManager, OstrisLinearLayerMemoryManager, _DEVICE_STATE, ) import random LINEAR_MODULES = [ "Linear", "LoRACompatibleLinear", "QLinear", 'OstrisLinear', ] CONV_MODULES = [ "Conv2d", "LoRACompatibleConv", "QConv2d", ] UNMANAGED_MODULES = [ "LayerNorm", "BatchNorm1d", "BatchNorm2d", "BatchNorm3d", "GroupNorm", "InstanceNorm1d", "InstanceNorm2d", "InstanceNorm3d", "Embedding", "EmbeddingBag", "RNNBase", "LSTM", "GRU", "RNN", "Conv3d" ] UNMANAGED_MODULES_INCLUDES = ["RotaryEmbedding", "Norm", "RotaryPosEmbed"] class MemoryManager: def __init__( self, module: torch.nn.Module, process_device: torch.device = torch.device("cpu"), ): self.module: torch.nn.Module = module self.process_device: torch.device = process_device self.unmanaged_modules: list[torch.nn.Module] = [] def memory_managed_to(self, *args, **kwargs): # first move all the unmanaged modules for module in self.unmanaged_modules: if isinstance(module, torch.Tensor): # Parameters and bare tensor buffers cannot move this way module.data = module.data.to(*args, **kwargs) else: module.to(*args, **kwargs) # check for a dtype argument dtype = None if "dtype" in kwargs: dtype = kwargs["dtype"] elif len(args) > 0: for i, arg in enumerate(args): if isinstance(arg, torch.dtype): dtype = arg break if dtype is not None: return self.module._mm_to(dtype=dtype) return self.module @classmethod def attach( cls, module: torch.nn.Module, device: torch.device, offload_percent: float = 1.0, ignore_modules: list[torch.nn.Module] = [] ): if hasattr(module, "_memory_manager"): # already attached return module._memory_manager = cls(module, device) # override the to method to handle memory management module._mm_to = module.to module.to = module._memory_manager.memory_managed_to # add ignore modules to unmanaged list for im in ignore_modules: module._memory_manager.unmanaged_modules.append(im) # count ignore modules as processed modules_processed = [x for x in ignore_modules] # attach to all modules for name, sub_module in module.named_modules(): for child_name, child_module in sub_module.named_modules(): if ( child_module.__class__.__name__ in LINEAR_MODULES and child_module not in modules_processed ): skip = False if offload_percent < 1.0: # randomly skip some modules if random.random() > offload_percent: skip = True if skip: module._memory_manager.unmanaged_modules.append(child_module) else: # linear; OstrisLinear bounces its quantized buffers instead # of a dequantized weight (module.weight is a property) if getattr(child_module, "is_ostris_quantized", False): OstrisLinearLayerMemoryManager.attach( child_module, module._memory_manager ) else: LinearLayerMemoryManager.attach( child_module, module._memory_manager ) # attach to ARA as well if hasattr(child_module, "ara_lora_ref"): ara = child_module.ara_lora_ref() if ara not in modules_processed: MemoryManager.attach( ara, device, ) modules_processed.append(child_module) elif ( child_module.__class__.__name__ in CONV_MODULES and child_module not in modules_processed ): skip = False if offload_percent < 1.0: # randomly skip some modules if random.random() > offload_percent: skip = True if skip: module._memory_manager.unmanaged_modules.append(child_module) else: # conv ConvLayerMemoryManager.attach( child_module, module._memory_manager ) # attach to ARA as well if hasattr(child_module, "ara_lora_ref"): ara = child_module.ara_lora_ref() if ara not in modules_processed: MemoryManager.attach( ara, device, ) modules_processed.append(ara) modules_processed.append(child_module) elif child_module.__class__.__name__ in UNMANAGED_MODULES or any( inc in child_module.__class__.__name__ for inc in UNMANAGED_MODULES_INCLUDES ): # unmanaged module._memory_manager.unmanaged_modules.append(child_module) else: continue @classmethod def detach(cls, module: torch.nn.Module): """ Reverse of attach(). Moves unmanaged modules back to CPU, restores the original .to() and forward methods on all child layers, unpins CPU weight tensors, and clears the global CUDA device state. Call this before unloading/replacing a module that had attach() applied. """ if not hasattr(module, "_memory_manager"): return for unmanaged in module._memory_manager.unmanaged_modules: try: if isinstance(unmanaged, torch.Tensor): unmanaged.data = unmanaged.data.to('cpu') else: unmanaged.to('cpu') except Exception: pass if hasattr(module, "_mm_to"): module.to = module._mm_to del module._mm_to del module._memory_manager for child in module.modules(): lmm = getattr(child, "_layer_memory_manager", None) if lmm is None: continue original_forward = getattr(lmm, "_original_forward", None) if original_forward is not None: if hasattr(child, "ara_lora_ref"): ara = child.ara_lora_ref() if ara is not None: ara.org_forward = original_forward else: child.forward = original_forward for param_name in ("weight", "bias"): # read _parameters directly: OstrisLinear.weight is a property that # materializes a full dequantized weight on access param = child._parameters.get(param_name, None) if param is None or not isinstance(param, torch.nn.Parameter): continue try: if param.data.is_pinned(): object.__setattr__( child, param_name, torch.nn.Parameter( param.data.clone(), requires_grad=param.requires_grad, ), ) except Exception: pass if getattr(child, "is_ostris_quantized", False): # move quantized buffers home and unpin them (clone drops pinning) for buf_name, buf in list(child._buffers.items()): if buf is None: continue try: if buf.device.type != "cpu": buf = buf.to("cpu") if buf.is_pinned(): buf = buf.clone() child._buffers[buf_name] = buf except Exception: pass del child._layer_memory_manager if hasattr(child, "_memory_management_device"): del child._memory_management_device if hasattr(child, "_is_memory_managed"): del child._is_memory_managed keys_to_delete = [ dev for dev in _DEVICE_STATE if isinstance(dev, torch.device) and dev.type == "cuda" ] for key in keys_to_delete: del _DEVICE_STATE[key] torch.cuda.empty_cache() @classmethod def free(cls, module: torch.nn.Module): """ Detach memory management (if attached) and destroy the module's weights by moving them to the meta device. Unlike detach(), nothing is staged back to CPU first: to('meta') frees each storage from wherever it currently lives, so no transient host allocation is made for data that is about to be discarded, and pinned tensors are freed without the clone that unpinning requires. Freed pinned storages land in torch's caching host allocator, not the OS; call release_cached_memory() afterward to get the RSS back. """ if hasattr(module, "_memory_manager"): if hasattr(module, "_mm_to"): module.to = module._mm_to del module._mm_to del module._memory_manager for child in module.modules(): lmm = getattr(child, "_layer_memory_manager", None) if lmm is None: continue original_forward = getattr(lmm, "_original_forward", None) if original_forward is not None: if hasattr(child, "ara_lora_ref"): ara = child.ara_lora_ref() if ara is not None: ara.org_forward = original_forward else: child.forward = original_forward del child._layer_memory_manager if hasattr(child, "_memory_management_device"): del child._memory_management_device if hasattr(child, "_is_memory_managed"): del child._is_memory_managed keys_to_delete = [ dev for dev in _DEVICE_STATE if isinstance(dev, torch.device) and dev.type == "cuda" ] for key in keys_to_delete: del _DEVICE_STATE[key] # bypass any overridden/nopped-out .to() so the storages are actually freed torch.nn.Module.to(module, "meta") torch.cuda.empty_cache() @classmethod def release_cached_memory(cls): """ Return freed memory to the OS. Freed pinned-host storages sit in torch's caching host allocator and freed pageable memory sits in glibc's arenas; neither shows up as reclaimed RSS without an explicit flush. Call after free()ing a large module. """ gc.collect() torch.cuda.empty_cache() # torch's pinned-host cache; private API, name varies by torch version for fn_name in ("_accelerator_emptyHostCache", "_host_emptyCache"): fn = getattr(torch._C, fn_name, None) if fn is not None: try: fn() break except Exception: pass # glibc keeps freed arenas mapped; CDLL(None) resolves malloc_trim in # the running process where glibc is present and fails cleanly on # macOS/musl try: ctypes.CDLL(None).malloc_trim(0) except Exception: pass