From 775977b7f3a633427fbdaea4a4e29c0e96c1dae9 Mon Sep 17 00:00:00 2001 From: Shubham Padkonde Date: Sun, 20 Sep 2026 08:24:36 +0530 Subject: [PATCH 1/3] Don't rewrite packed modules as float weights when resuming A run resumed through AR_RESUME_DIR skips the blocks a previous run already quantized, so those modules are still plain `nn.Linear` in the model tree while their packed tensors sit in the shards the crashed run flushed. `finalize()`'s remaining-weights pass deduplicates by exact tensor name, so `.weight` never matched the saved `.qweight` and the stale floating-point weight was written on top of the packed one, silently duplicating every module of every skipped block. Two changes: - Adopt the previous run's shards at the start of `finalize()`. Discovery previously ran on the first `_flush_shard()`, which `finalize()` only reaches after the remaining-weights pass, so that pass could not see anything the crashed run had written. - Skip a module's unpacked `weight` when its packed tensors are already saved, and log how many were skipped. Fixes #2350 Signed-off-by: Shubham Padkonde Co-Authored-By: Claude Opus 5 --- auto_round/compressors/shard_writer.py | 31 ++++++++++++++++ test/unit/common/utils/test_shard_writer.py | 39 +++++++++++++++++++++ 2 files changed, 70 insertions(+) diff --git a/auto_round/compressors/shard_writer.py b/auto_round/compressors/shard_writer.py index 6232b3d884..2c1a3acd41 100644 --- a/auto_round/compressors/shard_writer.py +++ b/auto_round/compressors/shard_writer.py @@ -38,6 +38,11 @@ DEFAULT_MAX_SHARD_SIZE = "5GB" +# Parameter names that only exist once a module has been packed for export, and +# that replace its floating-point ``weight``. +PACKED_WEIGHT_NAMES = frozenset({"qweight", "weight_packed"}) + + class ShardWriter: """ Handles shard-saving of model parameters to disk with memory management. @@ -418,25 +423,51 @@ def _offload_to_meta(self, saved_params): def finalize(self) -> None: """Saves remaining weights, renames files, and writes the index JSON.""" + # Adopt the shards of a run this one resumed from before deciding which + # weights are still missing, otherwise their tensors look unsaved here. + if envs.AR_RESUME_DIR and not self._existing_shards_discovered: + self._discover_existing_shards() + self._existing_shards_discovered = True + # 1. Capture remaining weights not yet saved full_sd = self.model.state_dict() tie_word_embeddings = False if hasattr(self.model, "config") and hasattr(self.model.config, "tie_word_embeddings"): tie_word_embeddings = self.model.config.tie_word_embeddings + # Modules whose packed tensors were already written, e.g. by a previous + # run this one resumed from. Such a module is skipped by tuning, so the + # model tree still holds its original floating-point weight under a name + # that never matches the packed one. + packed_layers = { + pname.rsplit(".", 1)[0] for pname in self._all_saved if pname.rsplit(".", 1)[-1] in PACKED_WEIGHT_NAMES + } + finalize_skipped_meta_tensors = [] + stale_unpacked_tensors = [] for pname, tensor in full_sd.items(): if pname in self._all_saved: continue if tensor.device.type == "meta": continue layer_name = ".".join(pname.split(".")[:-1]) + if pname.rsplit(".", 1)[-1] == "weight" and layer_name in packed_layers: + # Writing it would put the module in the checkpoint twice: once + # packed and once as the stale floating-point weight. + stale_unpacked_tensors.append(pname) + continue if self.lm_head_name is not None and layer_name == self.lm_head_name and tie_word_embeddings: lm_head_module = get_module(self.model, self.lm_head_name) lm_head_module.to("meta") # Must to meta, otherwise model's saver will dump it again continue self._add_tensor(pname, tensor.detach().to("cpu")) + if stale_unpacked_tensors: + logger.info( + f"Skipped {len(stale_unpacked_tensors)} unpacked weight(s) of already-packed modules, " + f"e.g. {stale_unpacked_tensors[:3]}." + ) + self._flush_shard() total_skipped = len(self.skipped_meta_tensors) + len(finalize_skipped_meta_tensors) diff --git a/test/unit/common/utils/test_shard_writer.py b/test/unit/common/utils/test_shard_writer.py index 7c7a50c40b..68ccc7bd46 100644 --- a/test/unit/common/utils/test_shard_writer.py +++ b/test/unit/common/utils/test_shard_writer.py @@ -197,3 +197,42 @@ def test_oversized_tensor_does_not_leave_tiny_preceding_shard(tmp_path, monkeypa assert writer.shard_counter == 1 assert set(writer.current_shard_tensors) == set() + + +def test_finalize_skips_unpacked_weight_of_resumed_packed_module(tmp_path, monkeypatch): + """A module packed before a crash must not be written again as fp weights. + + The resumed run skips tuning for that module, so the model tree still holds + the original ``nn.Linear``. Its ``weight`` name never matches the packed + ``qweight`` recovered from the crashed run's shard, so the name-exact dedup + alone would write both into the final checkpoint. + """ + from auto_round import envs + + # A shard flushed by the crashed run, holding the packed tensors. + torch.save( + { + "transformer_blocks.0.linear.qweight": torch.zeros(4, 1, dtype=torch.int32), + "transformer_blocks.0.linear.scales": torch.ones(4, 1), + "transformer_blocks.0.linear.bias": torch.zeros(4), + }, + os.path.join(tmp_path, "model-shard-00001.bin"), + ) + + monkeypatch.setattr(envs, "AR_RESUME_DIR", str(tmp_path)) + + model = _DiffusionStyleModel() + writer = _make_writer(model, str(tmp_path), monkeypatch) + writer.finalize() + + saved = {} + for name in os.listdir(tmp_path): + if name.endswith(".bin"): + saved.update(torch.load(os.path.join(tmp_path, name), map_location="cpu")) + + assert "transformer_blocks.0.linear.qweight" in saved + assert ( + "transformer_blocks.0.linear.weight" not in saved + ), "the packed module must not also be saved as a floating-point weight" + # Modules that were never packed are still saved. + assert "proj_out.weight" in saved From f31a1b2bd567e5ce94c2bedfb5df1676fdb6bce7 Mon Sep 17 00:00:00 2001 From: Shubham Padkonde Date: Mon, 28 Sep 2026 17:02:35 +0530 Subject: [PATCH 2/3] fix: recover finalized shards and transformed resume names Signed-off-by: Shubham Padkonde --- auto_round/compressors/shard_writer.py | 34 +++++++------- test/unit/common/utils/test_shard_writer.py | 50 ++++++++++++++++----- 2 files changed, 58 insertions(+), 26 deletions(-) diff --git a/auto_round/compressors/shard_writer.py b/auto_round/compressors/shard_writer.py index 2c1a3acd41..b59d733c72 100644 --- a/auto_round/compressors/shard_writer.py +++ b/auto_round/compressors/shard_writer.py @@ -129,22 +129,21 @@ def _discover_existing_shards(self) -> None: shard numbering instead of colliding with them, and ``finalize()``'s index covers tensors from both processes. - Only files still in the pre-``finalize()`` temp naming - (``model-shard-NNNNN.``) are considered: once ``finalize()`` runs - it renames everything to the final HF layout, so a directory with no - such temp files means either nothing has been flushed yet, or a prior - run already finished -- neither should be treated as in-progress - shards to adopt. + Include both temporary and final checkpoint names: the resume manifest + remains live until the export steps after ``finalize()`` succeed. """ output_dir = self.output_dir if not os.path.isdir(output_dir): return pattern = re.compile(rf"^model-shard-(\d+)\.{re.escape(self.shard_suffix)}$") + final_pattern = re.compile(rf"^model-(\d+)-of-\d+\.{re.escape(self.shard_suffix)}$") found = [] for fname in os.listdir(output_dir): - m = pattern.match(fname) + m = pattern.match(fname) or final_pattern.match(fname) if m: found.append((int(m.group(1)), fname)) + elif fname == f"model.{self.shard_suffix}": + found.append((1, fname)) if not found: return found.sort() @@ -153,7 +152,7 @@ def _discover_existing_shards(self) -> None: params = self._read_shard_tensor_names(path) self.shard_meta.append({"tmp_file": fname, "params": params, "dir": output_dir}) self._all_saved.update(params) - self.shard_counter = found[-1][0] + self.shard_counter = max(len(found), found[-1][0]) logger.info( f"ShardWriter: discovered {len(found)} already-flushed shard(s) in {output_dir} " f"from a previous run; resuming shard numbering from {self.shard_counter}." @@ -326,11 +325,9 @@ def _add_tensor(self, name: str, tensor: torch.Tensor): self._add_tensor(sub_name, sub_tensor) return - # transformers will handle _checkpoint_conversion_mapping automatically if is_immediate_saving=False - if self.reverse_weight_transforms is not None: - name = revert_name_with_weight_transforms(name, self.reverse_weight_transforms) - else: - name = revert_checkpoint_conversion_mapping(name, self.reverse_checkpoint_conversion_mapping) + name = self._checkpoint_name(name) + if name in self._all_saved or name in self.current_shard_tensors: + return t_size = tensor.nbytes self.total_param_elems += tensor.numel() @@ -351,6 +348,12 @@ def _add_tensor(self, name: str, tensor: torch.Tensor): self.current_shard_tensors[name] = tensor self.current_shard_size += t_size + def _checkpoint_name(self, name: str) -> str: + """Use the serialized namespace for both saving and resume comparisons.""" + if self.reverse_weight_transforms is not None: + return revert_name_with_weight_transforms(name, self.reverse_weight_transforms) + return revert_checkpoint_conversion_mapping(name, self.reverse_checkpoint_conversion_mapping) + def _handle_tied_weights(self): """ Detects tied weights in the current shard and ensures they are only saved once. @@ -446,12 +449,13 @@ def finalize(self) -> None: finalize_skipped_meta_tensors = [] stale_unpacked_tensors = [] for pname, tensor in full_sd.items(): - if pname in self._all_saved: + checkpoint_name = self._checkpoint_name(pname) + if pname in self._all_saved or checkpoint_name in self._all_saved: continue if tensor.device.type == "meta": continue layer_name = ".".join(pname.split(".")[:-1]) - if pname.rsplit(".", 1)[-1] == "weight" and layer_name in packed_layers: + if pname.rsplit(".", 1)[-1] == "weight" and checkpoint_name.rsplit(".", 1)[0] in packed_layers: # Writing it would put the module in the checkpoint twice: once # packed and once as the stale floating-point weight. stale_unpacked_tensors.append(pname) diff --git a/test/unit/common/utils/test_shard_writer.py b/test/unit/common/utils/test_shard_writer.py index 68ccc7bd46..b950cf5b74 100644 --- a/test/unit/common/utils/test_shard_writer.py +++ b/test/unit/common/utils/test_shard_writer.py @@ -15,6 +15,7 @@ import os from types import SimpleNamespace +import pytest import torch from auto_round.compressors.shard_writer import ShardWriter @@ -199,7 +200,12 @@ def test_oversized_tensor_does_not_leave_tiny_preceding_shard(tmp_path, monkeypa assert set(writer.current_shard_tensors) == set() -def test_finalize_skips_unpacked_weight_of_resumed_packed_module(tmp_path, monkeypatch): +@pytest.mark.parametrize("shard_name", ["model-shard-00001.bin", "model.bin", "model-00001-of-00002.bin"]) +@pytest.mark.parametrize("mapped", [False, True]) +@pytest.mark.parametrize("safe_serialization", [False, True]) +def test_finalize_skips_unpacked_weight_of_resumed_packed_module( + tmp_path, monkeypatch, shard_name, mapped, safe_serialization +): """A module packed before a crash must not be written again as fp weights. The resumed run skips tuning for that module, so the model tree still holds @@ -209,30 +215,52 @@ def test_finalize_skips_unpacked_weight_of_resumed_packed_module(tmp_path, monke """ from auto_round import envs - # A shard flushed by the crashed run, holding the packed tensors. - torch.save( + if safe_serialization: + from safetensors.torch import load_file, save_file + + save = save_file + load = load_file + shard_name = shard_name.replace(".bin", ".safetensors") + else: + save = torch.save + load = torch.load + + multiple_shards = "-of-" in shard_name + suffix = "safetensors" if safe_serialization else "bin" + prefix = "saved_blocks" if mapped else "transformer_blocks" + # A shard flushed by the crashed run, possibly already renamed by finalize. + save( { - "transformer_blocks.0.linear.qweight": torch.zeros(4, 1, dtype=torch.int32), - "transformer_blocks.0.linear.scales": torch.ones(4, 1), - "transformer_blocks.0.linear.bias": torch.zeros(4), + f"{prefix}.0.linear.qweight": torch.zeros(4, 1, dtype=torch.int32), + f"{prefix}.0.linear.scales": torch.ones(4, 1), + f"{prefix}.0.linear.bias": torch.full((4,), 7.0), }, - os.path.join(tmp_path, "model-shard-00001.bin"), + os.path.join(tmp_path, shard_name), ) + if multiple_shards: + save({"completed.qweight": torch.ones(2)}, str(tmp_path / f"model-00002-of-00002.{suffix}")) monkeypatch.setattr(envs, "AR_RESUME_DIR", str(tmp_path)) model = _DiffusionStyleModel() writer = _make_writer(model, str(tmp_path), monkeypatch) + writer.use_safetensors = safe_serialization + writer.shard_suffix = suffix + if mapped: + writer.reverse_checkpoint_conversion_mapping = {r"^transformer_blocks": "saved_blocks"} writer.finalize() saved = {} for name in os.listdir(tmp_path): - if name.endswith(".bin"): - saved.update(torch.load(os.path.join(tmp_path, name), map_location="cpu")) + if name.endswith(f".{suffix}"): + saved.update(load(os.path.join(tmp_path, name))) - assert "transformer_blocks.0.linear.qweight" in saved + assert f"{prefix}.0.linear.qweight" in saved assert ( - "transformer_blocks.0.linear.weight" not in saved + f"{prefix}.0.linear.weight" not in saved ), "the packed module must not also be saved as a floating-point weight" + assert torch.equal(saved[f"{prefix}.0.linear.bias"], torch.full((4,), 7.0)) + if multiple_shards: + assert torch.equal(saved["completed.qweight"], torch.ones(2)) # Modules that were never packed are still saved. assert "proj_out.weight" in saved From 007955ef5ee1ba8c702e378d0aaffd32838bdf0c Mon Sep 17 00:00:00 2001 From: Shubham Padkonde Date: Fri, 2 Oct 2026 21:59:35 +0530 Subject: [PATCH 3/3] fix: include adopted shards in checkpoint statistics Signed-off-by: Shubham Padkonde --- auto_round/compressors/shard_writer.py | 21 +++++++++++++++------ test/unit/common/utils/test_shard_writer.py | 18 ++++++++++++++++-- 2 files changed, 31 insertions(+), 8 deletions(-) diff --git a/auto_round/compressors/shard_writer.py b/auto_round/compressors/shard_writer.py index b59d733c72..e34b48b8dc 100644 --- a/auto_round/compressors/shard_writer.py +++ b/auto_round/compressors/shard_writer.py @@ -13,6 +13,7 @@ # limitations under the License. import json +import math import os import re from collections import OrderedDict @@ -149,26 +150,34 @@ def _discover_existing_shards(self) -> None: found.sort() for _, fname in found: path = os.path.join(output_dir, fname) - params = self._read_shard_tensor_names(path) + params, numel, size_bytes = self._read_shard_metadata(path) self.shard_meta.append({"tmp_file": fname, "params": params, "dir": output_dir}) self._all_saved.update(params) + self.total_param_elems += numel + self.total_param_size_bytes += size_bytes self.shard_counter = max(len(found), found[-1][0]) logger.info( f"ShardWriter: discovered {len(found)} already-flushed shard(s) in {output_dir} " f"from a previous run; resuming shard numbering from {self.shard_counter}." ) - def _read_shard_tensor_names(self, path: str) -> list[str]: - """Read only the tensor-name header of an already-flushed shard file, - without materializing any tensor data.""" + def _read_shard_metadata(self, path: str) -> tuple[list[str], int, int]: + """Read tensor names, element count and bytes without materializing data.""" if self.use_safetensors: from safetensors import safe_open with safe_open(path, framework="pt") as f: - return list(f.keys()) + params = list(f.keys()) + numel = sum(math.prod(f.get_slice(name).get_shape()) for name in params) + # A validated safetensors file stores all tensor bytes after its + # eight-byte header length and JSON header, with no gaps or padding. + with open(path, "rb") as f: + header_size = int.from_bytes(f.read(8), "little") + size_bytes = os.path.getsize(path) - 8 - header_size + return params, numel, size_bytes else: sd = torch.load(path, map_location="meta") - return list(sd.keys()) + return list(sd.keys()), sum(t.numel() for t in sd.values()), sum(t.nbytes for t in sd.values()) @property def output_dir(self) -> str: diff --git a/test/unit/common/utils/test_shard_writer.py b/test/unit/common/utils/test_shard_writer.py index b950cf5b74..d0aafccb4d 100644 --- a/test/unit/common/utils/test_shard_writer.py +++ b/test/unit/common/utils/test_shard_writer.py @@ -12,6 +12,7 @@ # See the License for the specific language governing permissions and # limitations under the License. +import json import os from types import SimpleNamespace @@ -238,7 +239,14 @@ def test_finalize_skips_unpacked_weight_of_resumed_packed_module( os.path.join(tmp_path, shard_name), ) if multiple_shards: - save({"completed.qweight": torch.ones(2)}, str(tmp_path / f"model-00002-of-00002.{suffix}")) + save( + { + "completed.qweight": torch.ones(2, dtype=torch.int16), + "completed.scale": torch.tensor(1.0, dtype=torch.float64), + "completed.empty": torch.empty(0, 3, dtype=torch.bfloat16), + }, + str(tmp_path / f"model-00002-of-00002.{suffix}"), + ) monkeypatch.setattr(envs, "AR_RESUME_DIR", str(tmp_path)) @@ -261,6 +269,12 @@ def test_finalize_skips_unpacked_weight_of_resumed_packed_module( ), "the packed module must not also be saved as a floating-point weight" assert torch.equal(saved[f"{prefix}.0.linear.bias"], torch.full((4,), 7.0)) if multiple_shards: - assert torch.equal(saved["completed.qweight"], torch.ones(2)) + assert torch.equal(saved["completed.qweight"], torch.ones(2, dtype=torch.int16)) # Modules that were never packed are still saved. assert "proj_out.weight" in saved + + # The index covers recovered packed tensors as well as newly written weights. + with open(tmp_path / f"model.{suffix}.index.json", encoding="utf-8") as index_file: + index = json.load(index_file) + assert index["metadata"]["total_parameters"] == sum(tensor.numel() for tensor in saved.values()) + assert index["metadata"]["total_size"] == sum(tensor.nbytes for tensor in saved.values())