[diffusion] feat: support TP for Flux.1.dev (#15666)
Co-authored-by: Mick <mickjagger19@icloud.com>
This commit is contained in:
@@ -6,7 +6,10 @@ import torch
|
|||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
|
|
||||||
from sglang.multimodal_gen.runtime.layers.activation import get_act_fn
|
from sglang.multimodal_gen.runtime.layers.activation import get_act_fn
|
||||||
from sglang.multimodal_gen.runtime.layers.linear import ReplicatedLinear
|
from sglang.multimodal_gen.runtime.layers.linear import (
|
||||||
|
ColumnParallelLinear,
|
||||||
|
RowParallelLinear,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class MLP(nn.Module):
|
class MLP(nn.Module):
|
||||||
@@ -25,18 +28,21 @@ class MLP(nn.Module):
|
|||||||
prefix: str = "",
|
prefix: str = "",
|
||||||
):
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.fc_in = ReplicatedLinear(
|
self.fc_in = ColumnParallelLinear(
|
||||||
input_dim,
|
input_dim,
|
||||||
mlp_hidden_dim, # For activation func like SiLU that need 2x width
|
mlp_hidden_dim,
|
||||||
bias=bias,
|
bias=True,
|
||||||
params_dtype=dtype,
|
gather_output=False,
|
||||||
)
|
)
|
||||||
|
|
||||||
self.act = get_act_fn(act_type)
|
self.act = get_act_fn(act_type)
|
||||||
if output_dim is None:
|
if output_dim is None:
|
||||||
output_dim = input_dim
|
output_dim = input_dim
|
||||||
self.fc_out = ReplicatedLinear(
|
self.fc_out = RowParallelLinear(
|
||||||
mlp_hidden_dim, output_dim, bias=bias, params_dtype=dtype
|
mlp_hidden_dim,
|
||||||
|
output_dim,
|
||||||
|
bias=True,
|
||||||
|
input_is_parallel=True,
|
||||||
)
|
)
|
||||||
|
|
||||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
|
|||||||
@@ -247,6 +247,7 @@ def load_model_from_full_model_state_dict(
|
|||||||
NotImplementedError: If got FSDP with more than 1D.
|
NotImplementedError: If got FSDP with more than 1D.
|
||||||
"""
|
"""
|
||||||
meta_sd = model.state_dict()
|
meta_sd = model.state_dict()
|
||||||
|
param_dict = dict(model.named_parameters())
|
||||||
sharded_sd = {}
|
sharded_sd = {}
|
||||||
custom_param_sd, reverse_param_names_mapping = hf_to_custom_state_dict(
|
custom_param_sd, reverse_param_names_mapping = hf_to_custom_state_dict(
|
||||||
full_sd_iterator, param_names_mapping
|
full_sd_iterator, param_names_mapping
|
||||||
@@ -259,7 +260,23 @@ def load_model_from_full_model_state_dict(
|
|||||||
)
|
)
|
||||||
if not hasattr(meta_sharded_param, "device_mesh"):
|
if not hasattr(meta_sharded_param, "device_mesh"):
|
||||||
full_tensor = full_tensor.to(device=device, dtype=param_dtype)
|
full_tensor = full_tensor.to(device=device, dtype=param_dtype)
|
||||||
# In cases where parts of the model aren't sharded, some parameters will be plain tensors
|
actual_param = param_dict.get(target_param_name)
|
||||||
|
weight_loader = (
|
||||||
|
getattr(actual_param, "weight_loader", None)
|
||||||
|
if actual_param is not None
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
if weight_loader is not None:
|
||||||
|
sharded_tensor = torch.empty_like(
|
||||||
|
meta_sharded_param, device=device, dtype=param_dtype
|
||||||
|
)
|
||||||
|
temp_param = nn.Parameter(sharded_tensor)
|
||||||
|
for attr in ["output_dim", "input_dim", "is_sharded_weight"]:
|
||||||
|
if hasattr(actual_param, attr):
|
||||||
|
setattr(temp_param, attr, getattr(actual_param, attr))
|
||||||
|
weight_loader(temp_param, full_tensor)
|
||||||
|
sharded_tensor = temp_param.data
|
||||||
|
else:
|
||||||
sharded_tensor = full_tensor
|
sharded_tensor = full_tensor
|
||||||
else:
|
else:
|
||||||
full_tensor = full_tensor.to(device=device, dtype=param_dtype)
|
full_tensor = full_tensor.to(device=device, dtype=param_dtype)
|
||||||
|
|||||||
@@ -36,7 +36,7 @@ from sglang.multimodal_gen.runtime.layers.attention import USPAttention
|
|||||||
|
|
||||||
# from sglang.multimodal_gen.runtime.layers.layernorm import LayerNorm as LayerNorm
|
# from sglang.multimodal_gen.runtime.layers.layernorm import LayerNorm as LayerNorm
|
||||||
from sglang.multimodal_gen.runtime.layers.layernorm import RMSNorm
|
from sglang.multimodal_gen.runtime.layers.layernorm import RMSNorm
|
||||||
from sglang.multimodal_gen.runtime.layers.linear import ReplicatedLinear
|
from sglang.multimodal_gen.runtime.layers.linear import ColumnParallelLinear
|
||||||
from sglang.multimodal_gen.runtime.layers.mlp import MLP
|
from sglang.multimodal_gen.runtime.layers.mlp import MLP
|
||||||
from sglang.multimodal_gen.runtime.layers.rotary_embedding import (
|
from sglang.multimodal_gen.runtime.layers.rotary_embedding import (
|
||||||
NDRotaryEmbedding,
|
NDRotaryEmbedding,
|
||||||
@@ -96,13 +96,16 @@ class FluxAttention(torch.nn.Module, AttentionModuleMixin):
|
|||||||
self.norm_q = RMSNorm(dim_head, eps=eps)
|
self.norm_q = RMSNorm(dim_head, eps=eps)
|
||||||
self.norm_k = RMSNorm(dim_head, eps=eps)
|
self.norm_k = RMSNorm(dim_head, eps=eps)
|
||||||
|
|
||||||
# Use ReplicatedLinear for fused QKV projections
|
self.to_qkv = ColumnParallelLinear(
|
||||||
self.to_qkv = ReplicatedLinear(query_dim, self.inner_dim * 3, bias=bias)
|
query_dim, self.inner_dim * 3, bias=bias, gather_output=True
|
||||||
|
)
|
||||||
|
|
||||||
if not self.pre_only:
|
if not self.pre_only:
|
||||||
self.to_out = torch.nn.ModuleList([])
|
self.to_out = torch.nn.ModuleList([])
|
||||||
self.to_out.append(
|
self.to_out.append(
|
||||||
ReplicatedLinear(self.inner_dim, self.out_dim, bias=out_bias)
|
ColumnParallelLinear(
|
||||||
|
self.inner_dim, self.out_dim, bias=out_bias, gather_output=True
|
||||||
|
)
|
||||||
)
|
)
|
||||||
if dropout != 0.0:
|
if dropout != 0.0:
|
||||||
self.to_out.append(torch.nn.Dropout(dropout))
|
self.to_out.append(torch.nn.Dropout(dropout))
|
||||||
@@ -110,11 +113,15 @@ class FluxAttention(torch.nn.Module, AttentionModuleMixin):
|
|||||||
if added_kv_proj_dim is not None:
|
if added_kv_proj_dim is not None:
|
||||||
self.norm_added_q = RMSNorm(dim_head, eps=eps)
|
self.norm_added_q = RMSNorm(dim_head, eps=eps)
|
||||||
self.norm_added_k = RMSNorm(dim_head, eps=eps)
|
self.norm_added_k = RMSNorm(dim_head, eps=eps)
|
||||||
# Use ReplicatedLinear for added (encoder) QKV projections
|
self.to_added_qkv = ColumnParallelLinear(
|
||||||
self.to_added_qkv = ReplicatedLinear(
|
added_kv_proj_dim,
|
||||||
added_kv_proj_dim, self.inner_dim * 3, bias=added_proj_bias
|
self.inner_dim * 3,
|
||||||
|
bias=added_proj_bias,
|
||||||
|
gather_output=True,
|
||||||
|
)
|
||||||
|
self.to_add_out = ColumnParallelLinear(
|
||||||
|
self.inner_dim, query_dim, bias=out_bias, gather_output=True
|
||||||
)
|
)
|
||||||
self.to_add_out = ReplicatedLinear(self.inner_dim, query_dim, bias=out_bias)
|
|
||||||
|
|
||||||
self.attn = USPAttention(
|
self.attn = USPAttention(
|
||||||
num_heads=num_heads,
|
num_heads=num_heads,
|
||||||
@@ -196,9 +203,13 @@ class FluxSingleTransformerBlock(nn.Module):
|
|||||||
self.mlp_hidden_dim = int(dim * mlp_ratio)
|
self.mlp_hidden_dim = int(dim * mlp_ratio)
|
||||||
|
|
||||||
self.norm = AdaLayerNormZeroSingle(dim)
|
self.norm = AdaLayerNormZeroSingle(dim)
|
||||||
self.proj_mlp = ReplicatedLinear(dim, self.mlp_hidden_dim)
|
self.proj_mlp = ColumnParallelLinear(
|
||||||
|
dim, self.mlp_hidden_dim, bias=True, gather_output=True
|
||||||
|
)
|
||||||
self.act_mlp = nn.GELU(approximate="tanh")
|
self.act_mlp = nn.GELU(approximate="tanh")
|
||||||
self.proj_out = ReplicatedLinear(dim + self.mlp_hidden_dim, dim)
|
self.proj_out = ColumnParallelLinear(
|
||||||
|
dim + self.mlp_hidden_dim, dim, bias=True, gather_output=True
|
||||||
|
)
|
||||||
|
|
||||||
self.attn = FluxAttention(
|
self.attn = FluxAttention(
|
||||||
query_dim=dim,
|
query_dim=dim,
|
||||||
@@ -408,10 +419,15 @@ class FluxTransformer2DModel(CachableDiT):
|
|||||||
pooled_projection_dim=self.config.pooled_projection_dim,
|
pooled_projection_dim=self.config.pooled_projection_dim,
|
||||||
)
|
)
|
||||||
|
|
||||||
self.context_embedder = ReplicatedLinear(
|
self.context_embedder = ColumnParallelLinear(
|
||||||
self.config.joint_attention_dim, self.inner_dim
|
self.config.joint_attention_dim,
|
||||||
|
self.inner_dim,
|
||||||
|
bias=True,
|
||||||
|
gather_output=True,
|
||||||
|
)
|
||||||
|
self.x_embedder = ColumnParallelLinear(
|
||||||
|
self.config.in_channels, self.inner_dim, bias=True, gather_output=True
|
||||||
)
|
)
|
||||||
self.x_embedder = ReplicatedLinear(self.config.in_channels, self.inner_dim)
|
|
||||||
self.transformer_blocks = nn.ModuleList(
|
self.transformer_blocks = nn.ModuleList(
|
||||||
[
|
[
|
||||||
FluxTransformerBlock(
|
FluxTransformerBlock(
|
||||||
@@ -437,10 +453,11 @@ class FluxTransformer2DModel(CachableDiT):
|
|||||||
self.norm_out = AdaLayerNormContinuous(
|
self.norm_out = AdaLayerNormContinuous(
|
||||||
self.inner_dim, self.inner_dim, elementwise_affine=False, eps=1e-6
|
self.inner_dim, self.inner_dim, elementwise_affine=False, eps=1e-6
|
||||||
)
|
)
|
||||||
self.proj_out = ReplicatedLinear(
|
self.proj_out = ColumnParallelLinear(
|
||||||
self.inner_dim,
|
self.inner_dim,
|
||||||
self.config.patch_size * self.config.patch_size * self.out_channels,
|
self.config.patch_size * self.config.patch_size * self.out_channels,
|
||||||
bias=True,
|
bias=True,
|
||||||
|
gather_output=True,
|
||||||
)
|
)
|
||||||
|
|
||||||
def forward(
|
def forward(
|
||||||
|
|||||||
Reference in New Issue
Block a user