forked from Karylab-cklius/vllm
remove unnecessary padding (#18)
Signed-off-by: Thien Tran <gau.nernst@yahoo.com.sg>
This commit is contained in:
@@ -56,8 +56,6 @@ class FlashInferCutlassMxfp8LinearKernel(Mxfp8LinearKernel):
|
||||
|
||||
input_shape = x.shape
|
||||
input_2d = x.view(-1, K)
|
||||
M_orig = input_2d.shape[0]
|
||||
|
||||
min_dim = 128
|
||||
|
||||
assert min_dim <= K, (
|
||||
@@ -72,11 +70,6 @@ class FlashInferCutlassMxfp8LinearKernel(Mxfp8LinearKernel):
|
||||
f"out_features is too small for mm_mxfp8."
|
||||
)
|
||||
|
||||
M_padded = ((M_orig + min_dim - 1) // min_dim) * min_dim
|
||||
if M_padded != M_orig:
|
||||
pad_rows = M_padded - M_orig
|
||||
input_2d = torch.nn.functional.pad(input_2d, (0, 0, 0, pad_rows))
|
||||
|
||||
input_mxfp8, input_scale = mxfp8_e4m3_quantize(
|
||||
input_2d, is_sf_swizzled_layout=True
|
||||
)
|
||||
@@ -93,9 +86,6 @@ class FlashInferCutlassMxfp8LinearKernel(Mxfp8LinearKernel):
|
||||
backend="cutlass",
|
||||
)
|
||||
|
||||
if M_padded != M_orig:
|
||||
output = output[:M_orig, :]
|
||||
|
||||
if bias is not None:
|
||||
output = output + bias
|
||||
|
||||
|
||||
Reference in New Issue
Block a user