[AMD] Fix Aiter RMSNorm layout handling (#23974)
This commit is contained in:
@@ -284,6 +284,15 @@ class RMSNorm(MultiPlatformOp):
|
|||||||
residual: Optional[torch.Tensor] = None,
|
residual: Optional[torch.Tensor] = None,
|
||||||
post_residual_addition: Optional[torch.Tensor] = None,
|
post_residual_addition: Optional[torch.Tensor] = None,
|
||||||
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
|
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
|
||||||
|
# Aiter's RMSNorm kernels expect 2D contiguous inputs. Keep the
|
||||||
|
# already-safe layout as a zero-copy path, and only normalize strided or
|
||||||
|
# higher-rank views such as Q/K slices from packed QKV projections.
|
||||||
|
needs_reshape = x.dim() != 2 and residual is None
|
||||||
|
if needs_reshape:
|
||||||
|
original_shape = x.shape
|
||||||
|
x = x.contiguous().reshape(-1, original_shape[-1])
|
||||||
|
elif not x.is_contiguous():
|
||||||
|
x = x.contiguous()
|
||||||
if residual is not None:
|
if residual is not None:
|
||||||
residual_out = torch.empty_like(x)
|
residual_out = torch.empty_like(x)
|
||||||
output = torch.empty_like(x)
|
output = torch.empty_like(x)
|
||||||
@@ -298,7 +307,10 @@ class RMSNorm(MultiPlatformOp):
|
|||||||
self.variance_epsilon,
|
self.variance_epsilon,
|
||||||
)
|
)
|
||||||
return output, residual_out
|
return output, residual_out
|
||||||
return rms_norm(x, self.weight.data, self.variance_epsilon)
|
output = rms_norm(x, self.weight.data, self.variance_epsilon)
|
||||||
|
if needs_reshape:
|
||||||
|
output = output.reshape(original_shape)
|
||||||
|
return output
|
||||||
|
|
||||||
def forward_hip(
|
def forward_hip(
|
||||||
self,
|
self,
|
||||||
|
|||||||
Reference in New Issue
Block a user