Fix inflated row pitch when a CP round-robin shard has a single row (#34142)
This commit is contained in:
@@ -196,7 +196,18 @@ def dsa_cp_round_robin_split_data(input_: Union[torch.Tensor, List]):
|
||||
return input_[indices]
|
||||
|
||||
# for torch device tensor
|
||||
return input_.view(-1, cp_size, *input_.shape[1:])[:, cp_rank].contiguous()
|
||||
shard = input_.view(-1, cp_size, *input_.shape[1:])[:, cp_rank]
|
||||
# .contiguous() is not sufficient here. When tokens == cp_size every rank's
|
||||
# shard has a single row, and a size-1 outer dimension imposes no contiguity
|
||||
# constraint, so is_contiguous() is True whatever stride(0) is and
|
||||
# .contiguous() becomes a no-op. The shard then keeps the cp_size-inflated
|
||||
# row pitch (cp_size * row_numel instead of row_numel), which any kernel that
|
||||
# takes its row pitch from stride(0) will read as an oversized tensor.
|
||||
# Compare the pitch against the parent's explicitly, so the copy happens
|
||||
# exactly when the shard really is strided -- and not at all for cp_size == 1.
|
||||
if shard.stride(0) != input_.stride(0):
|
||||
shard = shard.clone(memory_format=torch.contiguous_format)
|
||||
return shard
|
||||
|
||||
|
||||
def cal_padded_tokens(forward_batch: "ForwardBatch"):
|
||||
|
||||
Reference in New Issue
Block a user