Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
84 changes: 64 additions & 20 deletions auto_round/compressors/shard_writer.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,7 @@
# limitations under the License.

import json
import math
import os
import re
from collections import OrderedDict
Expand All @@ -38,6 +39,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.
Expand Down Expand Up @@ -124,47 +130,54 @@ 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.<ext>``) 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()
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.shard_counter = found[-1][0]
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:
Expand Down Expand Up @@ -321,11 +334,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()
Expand All @@ -346,6 +357,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.
Expand Down Expand Up @@ -418,25 +435,52 @@ 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
Comment thread
Shubham-Padkonde marked this conversation as resolved.

# 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:
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 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)
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)
Expand Down
81 changes: 81 additions & 0 deletions test/unit/common/utils/test_shard_writer.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,9 +12,11 @@
# See the License for the specific language governing permissions and
# limitations under the License.

import json
import os
from types import SimpleNamespace

import pytest
import torch

from auto_round.compressors.shard_writer import ShardWriter
Expand Down Expand Up @@ -197,3 +199,82 @@ 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()


@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
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

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(
{
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, shard_name),
)
if multiple_shards:
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))

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(f".{suffix}"):
saved.update(load(os.path.join(tmp_path, name)))

assert f"{prefix}.0.linear.qweight" in saved
assert (
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, 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())
Loading