flux2-dev-qb2 / code /models /tt_dit /layers /feedforward.py
stisiTT's picture
Add files using upload-large-folder tool
9aa90e0 verified
Raw History Blame Contribute Delete
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,
)