[diffusion] fix: fix h3 rank-local fsdp qkv loading (#34294)
This commit is contained in:
@@ -257,6 +257,43 @@ def _resolve_tp_shard_dim(
|
|||||||
return False, None
|
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(
|
def tp_local_shape(
|
||||||
sources: list[SafetensorsSource],
|
sources: list[SafetensorsSource],
|
||||||
shard_dim: int | None,
|
shard_dim: int | None,
|
||||||
@@ -421,6 +458,9 @@ def try_load_rank_local_fsdp_state_dict(
|
|||||||
return None
|
return None
|
||||||
sources_by_target, reverse_param_names_mapping = checkpoint_sources
|
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():
|
for target_param_name, sources in sources_by_target.items():
|
||||||
meta_param = meta_sd.get(target_param_name)
|
meta_param = meta_sd.get(target_param_name)
|
||||||
assembled_shape = assembled_source_shape(sources)
|
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):
|
if any(source.dtype in _QUANTIZED_SAFETENSORS_DTYPES for source in sources):
|
||||||
return None
|
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_param_sd: dict[str, torch.Tensor | LocalFSDPShard] = {}
|
||||||
local_bytes = 0
|
local_bytes = 0
|
||||||
with ExitStack() as stack:
|
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):
|
for target_param_name in sorted(sources_by_target):
|
||||||
meta_param = meta_sd[target_param_name]
|
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 = (
|
local_shape, global_offset = (
|
||||||
dist_tensor._utils.compute_local_shape_and_global_offset(
|
dist_tensor._utils.compute_local_shape_and_global_offset(
|
||||||
meta_param.shape,
|
meta_param.shape,
|
||||||
@@ -450,21 +513,21 @@ def try_load_rank_local_fsdp_state_dict(
|
|||||||
meta_param.placements,
|
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:
|
else:
|
||||||
tensor = read_rank_local_tensor(
|
local_shape = global_shape
|
||||||
|
global_offset = (0,) * meta_param.ndim
|
||||||
|
|
||||||
|
tensor = read_fsdp_rank_local_tensor(
|
||||||
sources_by_target[target_param_name],
|
sources_by_target[target_param_name],
|
||||||
handles,
|
handles,
|
||||||
tuple(meta_param.shape),
|
global_shape=global_shape,
|
||||||
(0,) * meta_param.ndim,
|
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_param_sd[target_param_name] = tensor
|
|
||||||
local_bytes += tensor.numel() * tensor.element_size()
|
local_bytes += tensor.numel() * tensor.element_size()
|
||||||
|
|
||||||
logger.info(
|
logger.info(
|
||||||
|
|||||||
@@ -588,6 +588,14 @@ class MiniMaxH3Attention(nn.Module):
|
|||||||
weight = self.qkv_proj.weight
|
weight = self.qkv_proj.weight
|
||||||
base_loader = weight.weight_loader
|
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:
|
def _weight_loader(param: torch.Tensor, loaded_weight: torch.Tensor) -> None:
|
||||||
# The grouped checkpoint layout is
|
# The grouped checkpoint layout is
|
||||||
# [num_query_groups, q_per_group + k + v] before splitting.
|
# [num_query_groups, q_per_group + k + v] before splitting.
|
||||||
@@ -602,18 +610,14 @@ class MiniMaxH3Attention(nn.Module):
|
|||||||
tp_size=self.tp_size,
|
tp_size=self.tp_size,
|
||||||
):
|
):
|
||||||
return
|
return
|
||||||
reordered = _reorder_grouped_qkv_to_qkv(
|
base_loader(param, _reorder_checkpoint_weight(loaded_weight))
|
||||||
loaded_weight,
|
|
||||||
num_query_groups=arch.num_attention_heads,
|
|
||||||
heads_per_group=1,
|
|
||||||
head_dim=arch.attention_head_dim,
|
|
||||||
)
|
|
||||||
base_loader(param, reordered)
|
|
||||||
|
|
||||||
if hasattr(weight, "_weight_loader"):
|
if hasattr(weight, "_weight_loader"):
|
||||||
weight._weight_loader = _weight_loader
|
weight._weight_loader = _weight_loader
|
||||||
else:
|
else:
|
||||||
weight.weight_loader = _weight_loader
|
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(
|
def forward(
|
||||||
self,
|
self,
|
||||||
|
|||||||
@@ -215,6 +215,30 @@ class TestRankLocalSafetensorsRead(unittest.TestCase):
|
|||||||
|
|
||||||
torch.testing.assert_close(tensor, torch.cat((first, second))[1:5])
|
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):
|
def test_reads_zero_sized_rank_local_shard(self):
|
||||||
with tempfile.TemporaryDirectory() as temp_dir:
|
with tempfile.TemporaryDirectory() as temp_dir:
|
||||||
file_path = str(Path(temp_dir) / "model.safetensors")
|
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
|
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():
|
for name, tensor in model.state_dict().items():
|
||||||
if name in expected_fp32:
|
if name in expected_fp32:
|
||||||
assert tensor.dtype == torch.float32, name
|
assert tensor.dtype == torch.float32, name
|
||||||
|
|||||||
Reference in New Issue
Block a user