|
| 1 | +# SPDX-FileCopyrightText: 2026 ModelCloud.ai |
| 2 | +# SPDX-FileCopyrightText: 2026 qubitium@modelcloud.ai |
| 3 | +# SPDX-License-Identifier: Apache-2.0 |
| 4 | +# Contact: qubitium@modelcloud.ai, x.com/qubitium |
| 5 | + |
| 6 | +from abc import ABC, abstractmethod |
| 7 | + |
| 8 | +import torch |
| 9 | +from typing import Dict, Type |
| 10 | +from defuser.logger import logger |
| 11 | +from dataclasses import dataclass |
| 12 | + |
| 13 | + |
| 14 | +class ReplacementModuleBase(ABC, torch.nn.Module): |
| 15 | + """ |
| 16 | + Abstract base class for module replacement during calibration phase. |
| 17 | +
|
| 18 | + Replacement modules replace original modules to ensure all components |
| 19 | + receive data for proper quantization statistics. |
| 20 | +
|
| 21 | + Subclasses must: |
| 22 | + 1. Implement `original_module_class()` to return the target module class name |
| 23 | + 2. Implement `__init__()` with signature: |
| 24 | + (self, original, config) |
| 25 | + """ |
| 26 | + |
| 27 | + # Registry: module class name -> replacement module class |
| 28 | + _replacement_registry: Dict[str, Type["ReplacementModuleBase"]] = {} |
| 29 | + |
| 30 | + def __init_subclass__(cls, **kwargs): |
| 31 | + """Automatically register subclasses in the replacement registry.""" |
| 32 | + super().__init_subclass__(**kwargs) |
| 33 | + |
| 34 | + # Only register if it's a concrete implementation (not ABC) |
| 35 | + if not getattr(cls, "__abstractmethods__", None): |
| 36 | + if cls.original_module_class() is None: |
| 37 | + raise TypeError( |
| 38 | + f"{cls.__name__} must implement 'original_module_class()' class method " |
| 39 | + "to return the name of the module class it replaces" |
| 40 | + ) |
| 41 | + |
| 42 | + if cls.original_module_class() in cls._replacement_registry: |
| 43 | + existing = cls._replacement_registry[cls.original_module_class()] |
| 44 | + raise ValueError( |
| 45 | + f"Module '{cls.original_module_class()}' already registered to " |
| 46 | + f"{existing.__name__}. Cannot register {cls.__name__}." |
| 47 | + ) |
| 48 | + |
| 49 | + cls._replacement_registry[cls.original_module_class()] = cls |
| 50 | + logger.trace(f"Registered {cls.__name__} for replacing {cls.original_module_class()}") |
| 51 | + |
| 52 | + def __init__(self, original: torch.nn.Module): |
| 53 | + super().__init__() |
| 54 | + _global_tracker.register_replacement( |
| 55 | + name=str(id(self)), |
| 56 | + original=original, |
| 57 | + replacement=self, |
| 58 | + ) |
| 59 | + self._materialized = False |
| 60 | + |
| 61 | + @classmethod |
| 62 | + def get_replacement_class(cls, module_class_name: str) -> Type["ReplacementModuleBase"]: |
| 63 | + """Get replacement class for a given module class name.""" |
| 64 | + return cls._replacement_registry.get(module_class_name) |
| 65 | + |
| 66 | + @classmethod |
| 67 | + def is_registered(cls, module_class_name: str) -> bool: |
| 68 | + """Check if a module class has a replacement implementation.""" |
| 69 | + return module_class_name in cls._replacement_registry |
| 70 | + |
| 71 | + @classmethod |
| 72 | + def is_to_be_replaced( |
| 73 | + cls, |
| 74 | + original: torch.nn.Module, |
| 75 | + ) -> bool: |
| 76 | + """Determine if the given module should be replaced. |
| 77 | +
|
| 78 | + Users can extend this method to add custom logic for replacement. |
| 79 | + """ |
| 80 | + return cls.is_registered(original.__class__.__name__) |
| 81 | + |
| 82 | + @classmethod |
| 83 | + def get_registered_modules(cls) -> list: |
| 84 | + """Get list of all registered module class names.""" |
| 85 | + return list(cls._replacement_registry.keys()) |
| 86 | + |
| 87 | + @classmethod |
| 88 | + @abstractmethod |
| 89 | + def original_module_class(cls) -> str: |
| 90 | + """Return the class name of the module this replaces.""" |
| 91 | + pass |
| 92 | + |
| 93 | + @classmethod |
| 94 | + @abstractmethod |
| 95 | + def from_original( |
| 96 | + cls, |
| 97 | + original: torch.nn.Module, |
| 98 | + config, |
| 99 | + ) -> "ReplacementModuleBase": |
| 100 | + """Create replacement module from original module.""" |
| 101 | + pass |
| 102 | + |
| 103 | + def materialize_weights(self): |
| 104 | + """Materialize weights if needed.""" |
| 105 | + if not self._materialized: |
| 106 | + self._materialize_weights() |
| 107 | + self.post_process_materialization() |
| 108 | + |
| 109 | + def _materialize_weights(self) -> None: |
| 110 | + """Materialize weights from the original module. |
| 111 | +
|
| 112 | + Subclasses should override this method to implement |
| 113 | + weight materialization logic. |
| 114 | + """ |
| 115 | + pass |
| 116 | + |
| 117 | + def release_original_module(self) -> None: |
| 118 | + """Release reference to the original module to free memory.""" |
| 119 | + # Release from global tracker |
| 120 | + _global_tracker.release_original(self) |
| 121 | + |
| 122 | + def _get_original_module(self) -> torch.nn.Module: |
| 123 | + """Get the original module associated with this replacement.""" |
| 124 | + return _global_tracker.get_original(self) |
| 125 | + |
| 126 | + def post_process_materialization(self) -> None: |
| 127 | + """Mark the replacement module as materialized.""" |
| 128 | + self._materialized = True |
| 129 | + self.release_original_module() |
| 130 | + |
| 131 | + |
| 132 | +@dataclass |
| 133 | +class ReplacedModuleInfo: |
| 134 | + original_module: torch.nn.Module |
| 135 | + replacement_module: ReplacementModuleBase |
| 136 | + |
| 137 | + |
| 138 | +class ModuleReplacementTracker: |
| 139 | + """Tracker to maintain mapping between replacement modules and their original modules. |
| 140 | +
|
| 141 | + This is a singleton class - only one instance can exist. |
| 142 | + """ |
| 143 | + |
| 144 | + _instance = None |
| 145 | + _initialized = False |
| 146 | + |
| 147 | + def __new__(cls): |
| 148 | + if cls._instance is None: |
| 149 | + cls._instance = super(ModuleReplacementTracker, cls).__new__(cls) |
| 150 | + return cls._instance |
| 151 | + |
| 152 | + def __init__(self): |
| 153 | + # Only initialize once |
| 154 | + if ModuleReplacementTracker._initialized: |
| 155 | + return |
| 156 | + |
| 157 | + # Map from replacement module id to original module |
| 158 | + self._replacement_to_original: Dict[int, torch.nn.Module] = {} |
| 159 | + # Map from module name to ReplacedModuleInfo |
| 160 | + self._name_to_info: Dict[str, ReplacedModuleInfo] = {} |
| 161 | + |
| 162 | + ModuleReplacementTracker._initialized = True |
| 163 | + |
| 164 | + @classmethod |
| 165 | + def get_instance(cls) -> "ModuleReplacementTracker": |
| 166 | + """Get the singleton instance of the tracker.""" |
| 167 | + if cls._instance is None: |
| 168 | + cls._instance = cls() |
| 169 | + return cls._instance |
| 170 | + |
| 171 | + def register_replacement(self, name: str, original: torch.nn.Module, replacement: ReplacementModuleBase) -> None: |
| 172 | + """Register a module replacement.""" |
| 173 | + self._replacement_to_original[id(replacement)] = original |
| 174 | + self._name_to_info[name] = ReplacedModuleInfo(original_module=original, replacement_module=replacement) |
| 175 | + logger.trace(f"Registered replacement for module: {name}") |
| 176 | + |
| 177 | + def get_original(self, replacement: ReplacementModuleBase) -> torch.nn.Module: |
| 178 | + """Get the original module for a given replacement module.""" |
| 179 | + return self._replacement_to_original.get(id(replacement)) |
| 180 | + |
| 181 | + def get_info_by_name(self, name: str) -> ReplacedModuleInfo: |
| 182 | + """Get replacement info by module name.""" |
| 183 | + return self._name_to_info.get(name) |
| 184 | + |
| 185 | + def release_original(self, replacement: ReplacementModuleBase) -> None: |
| 186 | + """Release the original module associated with a replacement module.""" |
| 187 | + replacement_id = id(replacement) |
| 188 | + if replacement_id in self._replacement_to_original: |
| 189 | + original = self._replacement_to_original[replacement_id] |
| 190 | + # Delete the original module to free memory |
| 191 | + del original |
| 192 | + del self._replacement_to_original[replacement_id] |
| 193 | + logger.trace(f"Released original module for replacement {replacement_id}") |
| 194 | + |
| 195 | + def release_all_originals(self) -> None: |
| 196 | + """Release all tracked original modules.""" |
| 197 | + count = len(self._replacement_to_original) |
| 198 | + if count > 0: |
| 199 | + self._replacement_to_original.clear() |
| 200 | + logger.debug(f"Released {count} original modules from tracker") |
| 201 | + |
| 202 | + def clear(self) -> None: |
| 203 | + """Clear all tracked information.""" |
| 204 | + self._replacement_to_original.clear() |
| 205 | + self._name_to_info.clear() |
| 206 | + logger.debug("Cleared module replacement tracker") |
| 207 | + |
| 208 | + |
| 209 | +_global_tracker = ModuleReplacementTracker() |
0 commit comments