Download code/models/tt_dit/layers/feedforward.py from stisiTT/flux2-dev-qb2: direct link, hf CLI and curl.
- Browser
- Download file 5.27 kB
-
https://huggingface.co/stisiTT/flux2-dev-qb2/resolve/main/code/models/tt_dit/layers/feedforward.py
- Command line
-
hf download hf://stisiTT/flux2-dev-qb2/code/models/tt_dit/layers/feedforward.py
-
curl -L -o feedforward.py https://huggingface.co/stisiTT/flux2-dev-qb2/resolve/main/code/models/tt_dit/layers/feedforward.py
5.27 kB
| # SPDX-FileCopyrightText: © 2025 Tenstorrent USA, Inc. | |
| # SPDX-License-Identifier: Apache-2.0 | |
| import ttnn | |
| from .linear import ColParallelLinear, Linear, LoRAColParallelLinear, LoRARowParallelLinear, RowParallelLinear | |
| from .module import Module | |
| class FeedForward(Module): | |
| """ | |
| Linear layer with replicated weights | |
| """ | |
| def __init__( | |
| self, | |
| dim: int, | |
| dim_out=None, | |
| mult: int = 4, | |
| activation_fn: str = "gelu", | |
| inner_dim=None, | |
| bias: bool = True, | |
| mesh_device=None, | |
| ): | |
| super().__init__() | |
| if inner_dim is None: | |
| inner_dim = int(dim * mult) | |
| dim_out = dim_out if dim_out is not None else dim | |
| self.mesh_device = mesh_device | |
| self.dim = dim | |
| self.dim_out = dim_out | |
| self.inner_dim = inner_dim | |
| self.activation_fn = activation_fn | |
| self.bias = bias | |
| self.ff1 = Linear(dim, inner_dim, bias=bias, mesh_device=mesh_device, activation_fn=activation_fn) | |
| self.ff2 = Linear(inner_dim, dim_out, bias=bias, mesh_device=mesh_device) | |
| def forward(self, x: ttnn.Tensor, compute_kernel_config=None) -> ttnn.Tensor: | |
| ff1_out = self.ff1(x, compute_kernel_config=compute_kernel_config) | |
| return self.ff2(ff1_out, compute_kernel_config=compute_kernel_config) | |
| class ParallelFeedForward(Module): | |
| """ | |
| Linear layer implementing megatron-style parallelism. | |
| """ | |
| def __init__( | |
| self, | |
| dim: int, | |
| dim_out=None, | |
| mult: int = 4, | |
| activation_fn: str = "gelu", | |
| inner_dim=None, | |
| bias: bool = True, | |
| mesh_device=None, | |
| mesh_axis=0, | |
| fsdp_mesh_axis=None, | |
| ccl_manager=None, | |
| lora_enabled: bool = False, | |
| ff1_dtype=ttnn.bfloat16, | |
| ff2_dtype=ttnn.bfloat16, | |
| activation_dtype=None, | |
| pin_output_bf16=False, | |
| ): | |
| super().__init__() | |
| if inner_dim is None: | |
| inner_dim = int(dim * mult) | |
| dim_out = dim_out if dim_out is not None else dim | |
| self.mesh_device = mesh_device | |
| self.dim = dim | |
| self.dim_out = dim_out | |
| self.inner_dim = inner_dim | |
| self.activation_fn = activation_fn | |
| self.bias = bias | |
| self.mesh_axis = mesh_axis | |
| self.fsdp_mesh_axis = fsdp_mesh_axis | |
| if self.fsdp_mesh_axis is not None: | |
| assert self.mesh_axis != self.fsdp_mesh_axis | |
| ColCls = LoRAColParallelLinear if lora_enabled else ColParallelLinear | |
| RowCls = LoRARowParallelLinear if lora_enabled else RowParallelLinear | |
| # ff1 is the ColParallel projection whose input crosses the fabric, so it carries the | |
| # activation cast + output pin; ff2 (RowParallel) only takes a weight dtype. | |
| self.ff1 = ColCls( | |
| dim, | |
| inner_dim, | |
| bias=bias, | |
| dtype=ff1_dtype, | |
| mesh_device=mesh_device, | |
| activation_fn=activation_fn, | |
| mesh_axis=mesh_axis, | |
| fsdp_mesh_axis=fsdp_mesh_axis, | |
| ccl_manager=ccl_manager, | |
| activation_dtype=activation_dtype, | |
| pin_output_bf16=pin_output_bf16, | |
| ) | |
| self.ff2 = RowCls( | |
| inner_dim, | |
| dim_out, | |
| bias=bias, | |
| dtype=ff2_dtype, | |
| mesh_device=mesh_device, | |
| mesh_axis=mesh_axis, | |
| fsdp_mesh_axis=fsdp_mesh_axis, | |
| ccl_manager=ccl_manager, | |
| ) | |
| def forward( | |
| self, x: ttnn.Tensor, compute_kernel_config=None, parallel_config=None, default_block_size=None | |
| ) -> ttnn.Tensor: | |
| """ | |
| Expects x to be replicated. | |
| Return output fractured on columns. | |
| `default_block_size` is forwarded to ff1 only, for callers that have measured block sizes for | |
| their ff1 shape; ff2 keeps the generic path. | |
| """ | |
| ff1_out = self.ff1( | |
| x, | |
| compute_kernel_config=compute_kernel_config, | |
| parallel_config=parallel_config, | |
| default_block_size=default_block_size, | |
| ) | |
| return self.ff2(ff1_out, compute_kernel_config=compute_kernel_config) | |
| def forward_fused_addcmul( | |
| self, | |
| x: ttnn.Tensor, | |
| addcmul_a: ttnn.Tensor, | |
| addcmul_b: ttnn.Tensor, | |
| scalar: float = 1.0, | |
| compute_kernel_config=None, | |
| parallel_config=None, | |
| default_block_size=None, | |
| core_grid=None, | |
| ) -> ttnn.Tensor: | |
| """Fused FFN forward with addcmul fused at the RS final write step. | |
| Computes: addcmul_a + scalar * ff2(ff1(x)) * addcmul_b | |
| Both addcmul_a and addcmul_b are already at their per-TP-device [D/tp] slice — | |
| no AllGather or scatter matmul is required. | |
| `default_block_size` is forwarded to ff1 only, as in `forward`. | |
| """ | |
| ff1_out = self.ff1( | |
| x, | |
| compute_kernel_config=compute_kernel_config, | |
| parallel_config=parallel_config, | |
| default_block_size=default_block_size, | |
| core_grid=core_grid, | |
| ) | |
| return self.ff2.forward_fused_addcmul( | |
| ff1_out, | |
| addcmul_a, | |
| addcmul_b, | |
| scalar=scalar, | |
| compute_kernel_config=compute_kernel_config, | |
| ) | |