# 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, )