diff --git a/obliteratus/evaluation/evaluator.py b/obliteratus/evaluation/evaluator.py index 92801da..e79d7c6 100644 --- a/obliteratus/evaluation/evaluator.py +++ b/obliteratus/evaluation/evaluator.py @@ -87,7 +87,6 @@ class Evaluator: desc="Evaluating PPL", ): batch_texts = ds[i : i + self.batch_size][self.text_column] - # Defensive filtering in case a custom dataset returns unexpected # values after selection or transformation. batch_texts = [ diff --git a/obliteratus/models/loader.py b/obliteratus/models/loader.py index 4bf9815..38a3e0c 100644 --- a/obliteratus/models/loader.py +++ b/obliteratus/models/loader.py @@ -366,7 +366,7 @@ class ModelHandle: for k, v in self._original_state.items(): target = current_state[k].device if k in current_state else None restored[k] = v.to(target) if target is not None else v - self.model.load_state_dict(restored) + self.model.load_state_dict(restored, strict=False) def cleanup(self): """Remove temporary offload directory if one was auto-created.""" @@ -622,26 +622,46 @@ def load_model( load_kwargs["offload_folder"] = _offload_dir logger.info(f"Auto-created offload folder: {_offload_dir}") - # Reserve GPU headroom for inference (KV cache, activations, generate()). - # Without this, device_map="auto" packs 100% of layers onto GPU, leaving - # no room for forward passes or generation on tight-memory setups. - if dev.is_cuda(): - max_memory = {} - for i in range(dev.device_count()): - total = torch.cuda.get_device_properties(i).total_memory - # Reserve 15% or 2 GiB (whichever is larger) for inference headroom - reserve = max(int(total * 0.15), 2 * 1024 ** 3) - usable = total - reserve - max_memory[i] = f"{usable // (1024 ** 2)}MiB" - # Allow overflow to CPU RAM, capped at 85% of physical memory - # to leave room for the OS, Python runtime, and serialization buffers. - total_ram, _ = dev._system_memory_gb() - cpu_budget_gb = int(total_ram * 0.85) - max_memory["cpu"] = f"{max(cpu_budget_gb, 4)}GiB" - load_kwargs["max_memory"] = max_memory - logger.info( - f"GPU memory budget: {', '.join(f'GPU{k}={v}' for k, v in max_memory.items() if k != 'cpu')}" - ) + # Skip max_memory when quantization shrinks the model enough to fit + if quantization in ("4bit", "8bit") and est_gb > 0 and gpu_gb > 0: + quant_factor = 4 if quantization == "4bit" else 2 + quant_est_gb = est_gb / quant_factor + if quant_est_gb < gpu_gb * 0.7: + logger.info( + f"Quantized estimate ({quant_est_gb:.1f} GB) fits GPU " + f"({gpu_gb:.0f} GB) — skipping max_memory constraint" + ) + else: + if dev.is_cuda(): + max_memory = {} + for i in range(dev.device_count()): + total = torch.cuda.get_device_properties(i).total_memory + reserve = max(int(total * 0.15), 2 * 1024 ** 3) + usable = total - reserve + max_memory[i] = f"{usable // (1024 ** 2)}MiB" + total_ram, _ = dev._system_memory_gb() + cpu_budget_gb = int(total_ram * 0.85) + max_memory["cpu"] = f"{max(cpu_budget_gb, 4)}GiB" + load_kwargs["max_memory"] = max_memory + logger.info( + f"GPU memory budget: {', '.join(f'GPU{k}={v}' for k, v in max_memory.items() if k != 'cpu')}" + ) + else: + # No quantization — use original max_memory logic + if dev.is_cuda(): + max_memory = {} + for i in range(dev.device_count()): + total = torch.cuda.get_device_properties(i).total_memory + reserve = max(int(total * 0.15), 2 * 1024 ** 3) + usable = total - reserve + max_memory[i] = f"{usable // (1024 ** 2)}MiB" + total_ram, _ = dev._system_memory_gb() + cpu_budget_gb = int(total_ram * 0.85) + max_memory["cpu"] = f"{max(cpu_budget_gb, 4)}GiB" + load_kwargs["max_memory"] = max_memory + logger.info( + f"GPU memory budget: {', '.join(f'GPU{k}={v}' for k, v in max_memory.items() if k != 'cpu')}" + ) try: model = model_cls.from_pretrained(**load_kwargs) @@ -731,12 +751,20 @@ def load_model( ) else: handle.snapshot() - elif gpu_gb > 0 and est_gb > 0 and est_gb > gpu_gb * 0.5: - logger.warning( - f"Auto-skipping state dict snapshot to save memory " - f"(model ~{est_gb:.0f} GB vs GPU {gpu_gb:.0f} GB). " - f"Use skip_snapshot=False to force." - ) + elif gpu_gb > 0 and est_gb > 0: + effective_gb = est_gb + if quantization == "4bit": + effective_gb = est_gb / 4 + elif quantization == "8bit": + effective_gb = est_gb / 2 + if effective_gb > gpu_gb * 0.5: + logger.warning( + f"Auto-skipping state dict snapshot to save memory " + f"(model ~{effective_gb:.0f} GB vs GPU {gpu_gb:.0f} GB). " + f"Use skip_snapshot=False to force." + ) + else: + handle.snapshot() else: handle.snapshot() diff --git a/obliteratus/runner.py b/obliteratus/runner.py index 3d8cf39..d3bc069 100644 --- a/obliteratus/runner.py +++ b/obliteratus/runner.py @@ -38,6 +38,7 @@ def run_study(config: StudyConfig) -> AblationReport: dtype=config.model.dtype, trust_remote_code=config.model.trust_remote_code, num_labels=config.model.num_labels, + quantization=getattr(config.model, "quantization", None), ) console.print(f" Architecture: {handle.architecture}") console.print(f" Layers: {handle.num_layers} Heads: {handle.num_heads}")