Fix sharding of quantized models with non-power-of-2 bits (#3006)
This commit is contained in:
@@ -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,
|
||||
|
||||
@@ -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}"
|
||||
|
||||
Reference in New Issue
Block a user