diff --git a/gptqmodel/looper/stage_layer.py b/gptqmodel/looper/stage_layer.py index cbf9f1ed3..c49eb5516 100644 --- a/gptqmodel/looper/stage_layer.py +++ b/gptqmodel/looper/stage_layer.py @@ -34,7 +34,7 @@ from ..looper.paroquant_processor import ParoQuantProcessor from ..looper.qqq_processor import QQQProcessor from ..utils.device import get_device, get_device_new -from ..utils.looper_helpers import normalize_device_like +from ..utils.looper_helpers import find_last_quantized_layer_index, normalize_device_like from ..utils.logger import live_renderables_suppressed, log_time_block, setup_logger from ..utils.model import find_modules, get_layer_name, get_module from ..utils.offload import offload_to_disk @@ -45,40 +45,6 @@ from .module_looper import ModuleLooper -def _find_last_quantized_layer_index( - looper: "ModuleLooper", - *, - layer_modules: List[List[str]], - layer_names: Optional[List[str]], - layer_count: int, -) -> Optional[int]: - """Return the highest layer index whose tracked modules are not all dynamically skipped.""" - if looper.gptq_model.quantize_config.lm_head or not layer_names: - return None - - layer_module_names = { - name.split("#", 1)[0] - for module_group in layer_modules - for name in module_group - if name - } - if not layer_module_names: - return None - - last_quantized_layer_index = -1 - for candidate_layer_index in range(layer_count): - layer_name = get_layer_name(layer_names, candidate_layer_index) - for module_name in layer_module_names: - module_full_name = f"{layer_name}.{module_name}" - # If at least one module in this layer is not dynamically excluded, - # the layer still needs forward/quantization work. - if looper.gptq_model.quantize_config.dynamic_get(layer_name=module_full_name) != False: - last_quantized_layer_index = candidate_layer_index - break - - return last_quantized_layer_index - - def _should_drain_finalize_futures_synchronously( looper: "ModuleLooper", *, @@ -401,8 +367,8 @@ def run_layer_stage( # Trailing layers whose tracked modules are all dynamically excluded never # need another forward or finalize pass, so the loop can stop once the # final eligible layer has been processed. - last_quantized_layer_index = _find_last_quantized_layer_index( - looper, + last_quantized_layer_index = find_last_quantized_layer_index( + looper.gptq_model.quantize_config, layer_modules=layer_modules, layer_names=layer_names, layer_count=layer_count, diff --git a/gptqmodel/looper/weight_only_looper.py b/gptqmodel/looper/weight_only_looper.py index b3a8c66c7..87b9938ac 100644 --- a/gptqmodel/looper/weight_only_looper.py +++ b/gptqmodel/looper/weight_only_looper.py @@ -35,7 +35,13 @@ from ..utils.device import get_device from ..utils.device_telemetry import emit_device_telemetry from ..utils.logger import log_time_block, setup_logger -from ..utils.looper_helpers import device_ctx, normalize_device_like, rehome_module_to_device, select_forward_devices +from ..utils.looper_helpers import ( + device_ctx, + find_last_quantized_layer_index, + normalize_device_like, + rehome_module_to_device, + select_forward_devices, +) from ..utils.model import ( find_modules, get_layer_name, @@ -756,9 +762,34 @@ def loop(self, **kwargs): ) try: + # Trailing layers whose tracked modules are all dynamically excluded never + # need another forward or finalize pass, so the loop can stop once the + # final eligible layer has been processed. + last_quantized_layer_index = find_last_quantized_layer_index( + self.gptq_model.quantize_config, + layer_modules=layer_modules, + layer_names=layer_names, + layer_count=layer_count, + ) + for layer_index in range(total_layers): is_lm_head_module = layer_index >= layer_count + if ( + not is_lm_head_module + and last_quantized_layer_index is not None + and layer_index > last_quantized_layer_index + ): + # The remaining layers are fully skipped by dynamic config, so + # avoid entering another layer-level quantization cycle. + log.debug( + "StageLayer: early stop at layer=%s, last_quantized_layer=%s", + layer_index, + last_quantized_layer_index, + ) + pb.close() + break + # Transformer blocks and lm_head follow the same weight-only # lifecycle, but lm_head is resolved from the root model. if is_lm_head_module: @@ -816,6 +847,12 @@ def loop(self, **kwargs): if named is None: continue + # Match ModuleLooper's processor lifecycle: dynamically + # excluded modules must not enter device scheduling or + # quantization work. + if self.processor.is_skipped(named): + continue + preferred_device = layer_strategy_device_map.get(module_name) if preferred_device is not None: # Weight-only has no SubsetPlan, so store the same diff --git a/gptqmodel/looper/weight_only_processor.py b/gptqmodel/looper/weight_only_processor.py index 28315a6f9..31b08abe7 100644 --- a/gptqmodel/looper/weight_only_processor.py +++ b/gptqmodel/looper/weight_only_processor.py @@ -68,6 +68,11 @@ def __init__( ) self.lock = threading.Lock() + def is_skipped(self, module: NamedModule) -> bool: + """Report whether dynamic configuration excludes this module.""" + + return self.qcfg.dynamic_get(layer_name=module.full_name) is False + @staticmethod def _uses_direct_pack(qcfg: RTNConfig | GGUFConfig | FP8Config | BitsAndBytesConfig) -> bool: """Returns whether the method packs directly from the original dense weights.""" diff --git a/gptqmodel/utils/looper_helpers.py b/gptqmodel/utils/looper_helpers.py index 703471687..f99fe6347 100644 --- a/gptqmodel/utils/looper_helpers.py +++ b/gptqmodel/utils/looper_helpers.py @@ -20,7 +20,7 @@ from ..utils.env import env_flag from ..utils.inspect import get_supported_kwargs from ..utils.logger import setup_logger -from ..utils.model import move_to, nested_move_to +from ..utils.model import get_layer_name, move_to, nested_move_to from ..utils.safe import ThreadSafe from ..utils.torch import ALL_DEVICES, CPU, HAS_NPU, torch_sync @@ -65,6 +65,7 @@ def torch_replicate( if TYPE_CHECKING: from ..looper.loop_processor import LoopProcessor from ..models._const import DEVICE + from ..quantization.config import BaseQuantizeConfig __all__ = [ @@ -74,10 +75,44 @@ def torch_replicate( "select_forward_devices", "normalize_device_like", "clone_module_for_devices", + "find_last_quantized_layer_index", "forward_batch_worker", ] +def find_last_quantized_layer_index( + quantize_config: "BaseQuantizeConfig", + *, + layer_modules: List[List[str]], + layer_names: Optional[List[str]], + layer_count: int, +) -> Optional[int]: + """Return the final layer containing a module not excluded by dynamic config.""" + + if quantize_config.lm_head or not layer_names: + return None + + layer_module_names = { + name.split("#", 1)[0] + for module_group in layer_modules + for name in module_group + if name + } + if not layer_module_names: + return None + + last_quantized_layer_index = -1 + for candidate_layer_index in range(layer_count): + layer_name = get_layer_name(layer_names, candidate_layer_index) + if any( + quantize_config.dynamic_get(layer_name=f"{layer_name}.{module_name}") is not False + for module_name in layer_module_names + ): + last_quantized_layer_index = candidate_layer_index + + return last_quantized_layer_index + + @contextmanager def device_ctx(dev: Optional[torch.device | "DEVICE"]): """Temporarily set the thread-local device for CUDA/XPU backends.""" diff --git a/tests/test_looper_helpers.py b/tests/test_looper_helpers.py index dadf4cc04..2ec780004 100644 --- a/tests/test_looper_helpers.py +++ b/tests/test_looper_helpers.py @@ -12,6 +12,22 @@ def _set_current_batch_index(self, batch_index): self.current_batch_index = batch_index +class _DynamicConfig: + lm_head = False + + def dynamic_get(self, *, layer_name): + return False if layer_name.startswith(("layers.1.", "layers.2.")) else None + + +def test_find_last_quantized_layer_index_uses_dynamic_exclusions(): + assert looper_helpers.find_last_quantized_layer_index( + _DynamicConfig(), + layer_modules=[["linear#capture_only"]], + layer_names=["layers.0", "layers.1", "layers.2"], + layer_count=3, + ) == 0 + + class _RequiresAttentionMask(torch.nn.Module): def __init__(self): super().__init__() diff --git a/tests/test_stage_modules.py b/tests/test_stage_modules.py index b2c0ee7ed..48d62fac3 100644 --- a/tests/test_stage_modules.py +++ b/tests/test_stage_modules.py @@ -973,6 +973,9 @@ def __init__(self): def pre_quantize(self, module): return module + def should_quantize_layer(self, *_args): + return True + def post_quantize(self, module): return module @@ -1043,7 +1046,7 @@ def create_named_modules(self, module, full, is_lm_head_module, layer_index, lay layers=[torch.nn.Linear(64, 64) for _ in range(3)], layer_modules=[["foo"]], planning_layer_modules=[["foo"]], - layers_prefix="model.layers", + layer_names=["model.layers.0", "model.layers.1", "model.layers.2"], fallback=True, shared_kv_cache_dict={}, pb=pb, diff --git a/tests/test_weight_only_looper.py b/tests/test_weight_only_looper.py index 3917d0b9b..902cd8bde 100644 --- a/tests/test_weight_only_looper.py +++ b/tests/test_weight_only_looper.py @@ -10,6 +10,8 @@ import gptqmodel.looper.weight_only_looper as weight_only_looper_module from gptqmodel.looper.weight_only_looper import WeightOnlyLooper +from gptqmodel.looper.weight_only_processor import WeightOnlyProcessor +from gptqmodel.looper.named_module import NamedModule from gptqmodel.quantization.config import RTNConfig, VramStrategy @@ -70,6 +72,9 @@ def pb(self, iterable, *, output_interval=None): def info(self, *_args, **_kwargs): return None + def debug(self, *_args, **_kwargs): + return None + class _TinyLayer(nn.Module): def __init__(self): @@ -164,6 +169,7 @@ def __init__(self, qcfg): self.finalized = [] self.finalize_called = False self.quant_devices = [] + self.skip_checks = [] def name(self): return "fake_weight_only" @@ -171,6 +177,10 @@ def name(self): def collect_memory_info(self, layer_index): self.memory_calls.append(layer_index) + def is_skipped(self, module): + self.skip_checks.append(module.full_name) + return self.qcfg.dynamic_get(layer_name=module.full_name) is False + def quantize_module(self, module, *, device=None): self.quantized.append(module.full_name) self.quant_devices.append(device) @@ -232,6 +242,123 @@ def test_weight_only_looper_reports_logbar_progress(monkeypatch): assert fake_logger.progresses[2].closed is True +def test_weight_only_processor_reports_dynamic_exclusions(): + qcfg = RTNConfig( + bits=4, + group_size=4, + offload_to_disk=False, + device="cpu", + dynamic={r"-:^layers\.0\.linear_b$": {}}, + ) + processor = WeightOnlyProcessor(tokenizer=None, qcfg=qcfg) + + included = NamedModule( + nn.Linear(4, 4, bias=False), + name="linear_a", + full_name="layers.0.linear_a", + layer_index=0, + ) + excluded = NamedModule( + nn.Linear(4, 4, bias=False), + name="linear_b", + full_name="layers.0.linear_b", + layer_index=0, + ) + + assert processor.is_skipped(included) is False + assert processor.is_skipped(excluded) is True + + +def test_weight_only_looper_skips_dynamic_exclusions_before_device_scheduling(monkeypatch): + qcfg = RTNConfig( + bits=4, + group_size=4, + offload_to_disk=False, + device="cuda:0", + dynamic={r"-:^layers\.0\.linear_b$": {}}, + ) + qcfg.lm_head = False + fake_logger = _FakeLogger() + processor = _FakeProcessor(qcfg) + model = _FakeQModel(qcfg) + model.model.layers = nn.ModuleList([_WideLayer()]) + model.simple_layer_modules = lambda **_kwargs: [["linear_a", "linear_b", "linear_c"]] + + devices = [torch.device("cuda:0"), torch.device("cuda:1")] + submitted_devices = [] + + def fake_submit(device, fn, *args, **kwargs): + submitted_devices.append(device) + future = Future() + future.set_result(fn(*args, **kwargs)) + return future + + monkeypatch.setattr(weight_only_looper_module, "log", fake_logger) + monkeypatch.setattr(weight_only_looper_module, "select_forward_devices", lambda _device: devices) + monkeypatch.setattr(weight_only_looper_module, "device_ctx", lambda _device: nullcontext()) + monkeypatch.setattr(weight_only_looper_module, "move_to", lambda obj, *, device, dtype=None: obj) + monkeypatch.setattr(weight_only_looper_module, "rehome_module_to_device", lambda *args, **kwargs: None) + monkeypatch.setattr(weight_only_looper_module.DEVICE_THREAD_POOL, "submit", fake_submit) + monkeypatch.setattr( + weight_only_looper_module, + "get_layers_with_prefixes", + lambda _model, _nodes: (list(model.model.layers), ["layers.0"]), + ) + + looper = WeightOnlyLooper(model=model, processor=processor) + looper.loop() + + expected_devices = [torch.device("cuda:0"), torch.device("cuda:1")] + assert processor.skip_checks == [ + "layers.0.linear_a", + "layers.0.linear_b", + "layers.0.linear_c", + ] + assert processor.quantized == ["layers.0.linear_a", "layers.0.linear_c"] + assert processor.quant_devices == expected_devices + assert processor.finalized == [ + ("layers.0.linear_a", torch.device("cuda:0")), + ("layers.0.linear_c", torch.device("cuda:1")), + ] + assert submitted_devices == expected_devices + expected_devices + + +def test_weight_only_looper_fast_mode_stops_before_trailing_excluded_layers(monkeypatch): + qcfg = RTNConfig( + bits=4, + group_size=4, + offload_to_disk=False, + device="cpu", + dynamic={ + r"-:^layers\.1\.": {}, + r"-:^layers\.2\.": {}, + }, + ) + qcfg.lm_head = False + fake_logger = _FakeLogger() + processor = _FakeProcessor(qcfg) + model = _FakeQModel(qcfg) + model.model.layers = nn.ModuleList([_TinyLayer(), _TinyLayer(), _TinyLayer()]) + + monkeypatch.setattr(weight_only_looper_module, "log", fake_logger) + monkeypatch.setattr( + weight_only_looper_module, + "get_layers_with_prefixes", + lambda _model, _nodes: ( + list(model.model.layers), + ["layers.0", "layers.1", "layers.2"], + ), + ) + + looper = WeightOnlyLooper(model=model, processor=processor) + looper.loop() + + assert processor.memory_calls == [0] + assert processor.skip_checks == ["layers.0.linear"] + assert processor.quantized == ["layers.0.linear"] + assert processor.finalized == [("layers.0.linear", torch.device("cpu"))] + + def test_weight_only_looper_quantizes_subset_across_multiple_devices(monkeypatch): qcfg = RTNConfig(bits=4, group_size=4, offload_to_disk=False, device="cuda:0") qcfg.lm_head = False