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: - (
| 1743 | |
| 1744 | |
| 1745 | class 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 |
no outgoing calls