[loader] support presharded fused mlp loading (#19519)
This commit is contained in:
@@ -575,8 +575,13 @@ class MergedColumnParallelLinear(ColumnParallelLinear):
|
|||||||
current_shard_offset = 0
|
current_shard_offset = 0
|
||||||
shard_offsets: List[Tuple[int, int, int]] = []
|
shard_offsets: List[Tuple[int, int, int]] = []
|
||||||
for i, output_size in enumerate(self.output_sizes):
|
for i, output_size in enumerate(self.output_sizes):
|
||||||
shard_offsets.append((i, current_shard_offset, output_size))
|
effective_size = (
|
||||||
current_shard_offset += output_size
|
output_size // self.tp_size
|
||||||
|
if self.use_presharded_weights
|
||||||
|
else output_size
|
||||||
|
)
|
||||||
|
shard_offsets.append((i, current_shard_offset, effective_size))
|
||||||
|
current_shard_offset += effective_size
|
||||||
packed_dim = getattr(param, "packed_dim", None)
|
packed_dim = getattr(param, "packed_dim", None)
|
||||||
|
|
||||||
use_bitsandbytes_4bit = getattr(param, "use_bitsandbytes_4bit", False)
|
use_bitsandbytes_4bit = getattr(param, "use_bitsandbytes_4bit", False)
|
||||||
|
|||||||
Reference in New Issue
Block a user