[diffusion] fix: fix h3 rank-local fsdp qkv loading (#34294)

This commit is contained in:
Mick
2026-08-10 23:59:00 +08:00
committed by GitHub
parent d07ac32d05
commit ec9babe36c
4 changed files with 115 additions and 22 deletions
@@ -257,6 +257,43 @@ def _resolve_tp_shard_dim(
return False, None
def read_fsdp_rank_local_tensor(
sources: list[SafetensorsSource],
handles: dict[str, Any],
global_shape: tuple[int, ...],
local_shape: tuple[int, ...],
global_offset: tuple[int, ...],
transform: Callable[[torch.Tensor], torch.Tensor] | None,
) -> torch.Tensor:
if transform is None:
return read_rank_local_tensor(
sources,
handles,
local_shape,
global_offset,
)
loaded_tensor = read_rank_local_tensor(
sources,
handles,
global_shape,
(0,) * len(global_shape),
)
transformed_tensor = transform(loaded_tensor)
if tuple(transformed_tensor.shape) != global_shape:
raise RuntimeError(
"Rank-local checkpoint transform shape mismatch: "
f"transformed={tuple(transformed_tensor.shape)}, expected={global_shape}"
)
slices = tuple(
slice(offset, offset + size)
for offset, size in zip(global_offset, local_shape, strict=True)
)
# clone so the local shard does not retain the full transformed storage
return transformed_tensor[slices].clone(memory_format=torch.contiguous_format)
def tp_local_shape(
sources: list[SafetensorsSource],
shard_dim: int | None,
@@ -421,6 +458,9 @@ def try_load_rank_local_fsdp_state_dict(
return None
sources_by_target, reverse_param_names_mapping = checkpoint_sources
param_dict = dict(model.named_parameters())
rank_local_transforms: dict[str, Callable[[torch.Tensor], torch.Tensor]] = {}
for target_param_name, sources in sources_by_target.items():
meta_param = meta_sd.get(target_param_name)
assembled_shape = assembled_source_shape(sources)
@@ -431,6 +471,27 @@ def try_load_rank_local_fsdp_state_dict(
if any(source.dtype in _QUANTIZED_SAFETENSORS_DTYPES for source in sources):
return None
actual_param = get_param_for_weight_loading(
model,
param_dict,
target_param_name,
)
if actual_param is None:
continue
rank_local_transform = actual_param.__dict__.get("rank_local_weight_transform")
if rank_local_transform is not None:
rank_local_transforms[target_param_name] = rank_local_transform
continue
supported, _ = _resolve_tp_shard_dim(actual_param)
if not supported:
logger.info(
"Falling back from rank-local FSDP checkpoint loading for "
"custom weight loader: %s",
target_param_name,
)
return None
local_param_sd: dict[str, torch.Tensor | LocalFSDPShard] = {}
local_bytes = 0
with ExitStack() as stack:
@@ -442,7 +503,9 @@ def try_load_rank_local_fsdp_state_dict(
}
for target_param_name in sorted(sources_by_target):
meta_param = meta_sd[target_param_name]
if isinstance(meta_param, dist_tensor.DTensor):
global_shape = tuple(meta_param.shape)
is_sharded = isinstance(meta_param, dist_tensor.DTensor)
if is_sharded:
local_shape, global_offset = (
dist_tensor._utils.compute_local_shape_and_global_offset(
meta_param.shape,
@@ -450,21 +513,21 @@ def try_load_rank_local_fsdp_state_dict(
meta_param.placements,
)
)
tensor = read_rank_local_tensor(
sources_by_target[target_param_name],
handles,
tuple(local_shape),
tuple(global_offset),
)
local_param_sd[target_param_name] = LocalFSDPShard(tensor)
else:
tensor = read_rank_local_tensor(
sources_by_target[target_param_name],
handles,
tuple(meta_param.shape),
(0,) * meta_param.ndim,
)
local_param_sd[target_param_name] = tensor
local_shape = global_shape
global_offset = (0,) * meta_param.ndim
tensor = read_fsdp_rank_local_tensor(
sources_by_target[target_param_name],
handles,
global_shape=global_shape,
local_shape=tuple(local_shape),
global_offset=tuple(global_offset),
transform=rank_local_transforms.get(target_param_name),
)
local_param_sd[target_param_name] = (
LocalFSDPShard(tensor) if is_sharded else tensor
)
local_bytes += tensor.numel() * tensor.element_size()
logger.info(
@@ -588,6 +588,14 @@ class MiniMaxH3Attention(nn.Module):
weight = self.qkv_proj.weight
base_loader = weight.weight_loader
def _reorder_checkpoint_weight(loaded_weight: torch.Tensor) -> torch.Tensor:
return _reorder_grouped_qkv_to_qkv(
loaded_weight,
num_query_groups=arch.num_attention_heads,
heads_per_group=1,
head_dim=arch.attention_head_dim,
)
def _weight_loader(param: torch.Tensor, loaded_weight: torch.Tensor) -> None:
# The grouped checkpoint layout is
# [num_query_groups, q_per_group + k + v] before splitting.
@@ -602,18 +610,14 @@ class MiniMaxH3Attention(nn.Module):
tp_size=self.tp_size,
):
return
reordered = _reorder_grouped_qkv_to_qkv(
loaded_weight,
num_query_groups=arch.num_attention_heads,
heads_per_group=1,
head_dim=arch.attention_head_dim,
)
base_loader(param, reordered)
base_loader(param, _reorder_checkpoint_weight(loaded_weight))
if hasattr(weight, "_weight_loader"):
weight._weight_loader = _weight_loader
else:
weight.weight_loader = _weight_loader
# rank-local FSDP must reorder grouped QKV before selecting each shard
weight.rank_local_weight_transform = _reorder_checkpoint_weight
def forward(
self,
@@ -215,6 +215,30 @@ class TestRankLocalSafetensorsRead(unittest.TestCase):
torch.testing.assert_close(tensor, torch.cat((first, second))[1:5])
def test_applies_layout_transform_before_rank_local_slice(self):
with tempfile.TemporaryDirectory() as temp_dir:
file_path = str(Path(temp_dir) / "model.safetensors")
grouped_qkv = torch.arange(12, dtype=torch.bfloat16).reshape(6, 2)
save_file({"weight": grouped_qkv}, file_path)
def reorder_qkv(weight: torch.Tensor) -> torch.Tensor:
grouped = weight.view(2, 3, 2)
return grouped.permute(1, 0, 2).reshape(6, 2)
with safe_open(file_path, framework="pt", device="cpu") as handle:
tensor = rank_local_checkpoint.read_fsdp_rank_local_tensor(
[self._source(file_path, "weight", (6, 2))],
{file_path: handle},
global_shape=(6, 2),
local_shape=(2, 2),
global_offset=(2, 0),
transform=reorder_qkv,
)
expected = reorder_qkv(grouped_qkv)[2:4]
torch.testing.assert_close(tensor, expected)
self.assertEqual(tensor.untyped_storage().nbytes(), expected.nbytes)
def test_reads_zero_sized_rank_local_shard(self):
with tempfile.TemporaryDirectory() as temp_dir:
file_path = str(Path(temp_dir) / "model.safetensors")
@@ -164,6 +164,8 @@ def test_meta_model_enforces_mixed_precision_weight_contract():
)
assert model._fsdp_mixed_dtype_params
qkv_weight = model.blocks[0].attn.qkv_proj.weight
assert callable(qkv_weight.rank_local_weight_transform)
for name, tensor in model.state_dict().items():
if name in expected_fp32:
assert tensor.dtype == torch.float32, name