Fix sharding of quantized models with non-power-of-2 bits (#3006)

This commit is contained in:
Tarjei Mandt
2026-01-18 07:21:56 -08:00
committed by GitHub
parent d2bef3c6bb
commit ca14d3d835
3 changed files with 18 additions and 6 deletions
+4 -4
View File
@@ -423,7 +423,7 @@ class QuantizedAllToShardedLinear(Module):
def _extra_repr(self) -> str:
out_dims, in_dims = self.weight.shape
in_dims *= 32 // self.bits
in_dims = (in_dims * 32) // self.bits
out_dims *= self.group.size()
return (
f"input_dims={in_dims}, output_dims={out_dims}, bias={'bias' in self}, "
@@ -457,7 +457,7 @@ class QuantizedAllToShardedLinear(Module):
):
group = group or mx.distributed.init()
output_dims, input_dims = quantized_linear_layer.weight.shape
input_dims *= 32 // quantized_linear_layer.bits
input_dims = (input_dims * 32) // quantized_linear_layer.bits
sl = cls(
input_dims,
@@ -549,7 +549,7 @@ class QuantizedShardedToAllLinear(Module):
def _extra_repr(self) -> str:
out_dims, in_dims = self.weight.shape
in_dims *= (32 // self.bits) * self.group.size()
in_dims = (in_dims * 32) // self.bits * self.group.size()
return (
f"input_dims={in_dims}, output_dims={out_dims}, bias={'bias' in self}, "
f"group_size={self.group_size}, bits={self.bits}"
@@ -580,7 +580,7 @@ class QuantizedShardedToAllLinear(Module):
):
group = group or mx.distributed.init()
output_dims, input_dims = quantized_linear_layer.weight.shape
input_dims *= 32 // quantized_linear_layer.bits
input_dims = (input_dims * 32) // quantized_linear_layer.bits
sl = cls(
input_dims,
+2 -2
View File
@@ -233,7 +233,7 @@ class QuantizedLinear(Module):
def _extra_repr(self):
out_dims, in_dims = self.weight.shape
in_dims *= 32 // self.bits
in_dims = (in_dims * 32) // self.bits
return (
f"input_dims={in_dims}, output_dims={out_dims}, bias={'bias' in self}, "
f"group_size={self.group_size}, bits={self.bits}, mode={self.mode}"
@@ -338,7 +338,7 @@ class QQLinear(Module):
def _extra_repr(self):
out_dims, in_dims = self.weight.shape
if self.weight.dtype == mx.uint32:
in_dims *= 32 // self.bits
in_dims = (in_dims * 32) // self.bits
return (
f"input_dims={in_dims}, output_dims={out_dims}, "
f"group_size={self.group_size}, bits={self.bits}, mode={self.mode}"