MCPcopy Create free account
hub / github.com/deepspeedai/DeepSpeed / SubParamLinearLayer

Class SubParamLinearLayer

deepspeed/module_inject/layers.py:1745–1880  ·  view source on GitHub ↗

Column-parallel linear layer with sub-parameter support. Handles cases where weights contain multiple logical sub-parameters that need to be partitioned separately (e.g., fused QKV, chunked MLP, GQA). The `shape` parameter controls how the weight is viewed and partitioned: - (

Source from the content-addressed store, hash-verified

1743
1744
1745class SubParamLinearLayer(TensorParallel_Layer):
1746 """
1747 Column-parallel linear layer with sub-parameter support.
1748
1749 Handles cases where weights contain multiple logical sub-parameters
1750 that need to be partitioned separately (e.g., fused QKV, chunked MLP, GQA).
1751
1752 The `shape` parameter controls how the weight is viewed and partitioned:
1753 - (3, -1) with partition_dim=0: 3 equal sub-params, partition each at dim 0
1754 - ((q, k, v), -1) with partition_dim=0: 3 unequal sub-params (1-level nesting)
1755 """
1756
1757 def __init__(self, module, mp_group, shape, partition_dim=0, **kwargs):
1758 super(SubParamLinearLayer, self).__init__(mp_group, **kwargs)
1759 self.weight = module.weight
1760 self.bias = module.bias
1761 self.shape = shape
1762 self.partition_dim = partition_dim
1763
1764 self._orig_weight_shape = tuple(module.weight.shape)
1765 self._orig_bias_shape = tuple(module.bias.shape) if self.bias is not None else None
1766 (self._logical_shape, self._output_shape, self._subparam_sizes,
1767 self._bias_partition_dim) = _infer_subparam_logical_shapes(self._orig_weight_shape, self.shape,
1768 self.partition_dim, self.name)
1769 # Resolve the per-rank widths once, for the same reason _freeze_partition_sizes does:
1770 # the split depends on this model's tp_meta.
1771 self._subparam_shard_widths = _subparam_shard_widths(
1772 self._subparam_sizes or (self._logical_shape[self.partition_dim], ), self.tp_world_size, self.tp_meta,
1773 self.name)
1774 self._bias_shape_spec = _bias_subparam_shape_spec(self._output_shape, self._bias_partition_dim,
1775 self._subparam_sizes)
1776 if self.bias is not None and self.bias.numel() != _shape_prod(self._output_shape):
1777 raise ValueError(f"AutoTP layer '{self.name}' bias size {self.bias.numel()} does not match output shape "
1778 f"{self._output_shape}.")
1779
1780 if self._should_materialize_tp_partition():
1781 self._tp_partition([self.weight, self.bias])
1782 self.support_training = True
1783 self.config_tp_params(self.weight)
1784 if self.bias is not None:
1785 self.config_tp_params(self.bias)
1786 self._mark_uc_metadata()
1787
1788 def forward(self, input):
1789 self._assert_compiled_if_deferred()
1790 if getattr(self, 'mp_group', None) is not None and not self.defer_collectives_to_compiler:
1791 input = ColumnParallel.apply(self.mp_group, input)
1792 output = torch.matmul(input, self.weight.transpose(-1, -2))
1793 if self.bias is not None:
1794 output = add_bias(output, self.bias)
1795 return output
1796
1797 @torch.no_grad()
1798 def gather_params(self, params_list):
1799 """Gather partitioned parameters back to full size."""
1800 for idx, param in enumerate(params_list):
1801 if param is None:
1802 continue

Calls

no outgoing calls