[diffusion] feat: allow T5's TP Group to reuse the transformer's SP Group (#17818)
Co-authored-by: Mick <mickjagger19@icloud.com>
This commit is contained in:
@@ -81,6 +81,11 @@ class EncoderConfig(ModelConfig):
|
|||||||
class TextEncoderConfig(EncoderConfig):
|
class TextEncoderConfig(EncoderConfig):
|
||||||
arch_config: ArchConfig = field(default_factory=TextEncoderArchConfig)
|
arch_config: ArchConfig = field(default_factory=TextEncoderArchConfig)
|
||||||
|
|
||||||
|
# Use the SP Group of the transformer as the TP Group of T5.
|
||||||
|
parallel_folding: bool = False
|
||||||
|
# "sp" or "ulysses" or "ring"
|
||||||
|
parallel_folding_mode: str = "sp"
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class ImageEncoderConfig(EncoderConfig):
|
class ImageEncoderConfig(EncoderConfig):
|
||||||
|
|||||||
@@ -1,12 +1,14 @@
|
|||||||
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
|
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
|
||||||
|
|
||||||
# SPDX-License-Identifier: Apache-2.0
|
# SPDX-License-Identifier: Apache-2.0
|
||||||
|
import argparse
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
|
|
||||||
from sglang.multimodal_gen.configs.models.encoders.base import (
|
from sglang.multimodal_gen.configs.models.encoders.base import (
|
||||||
TextEncoderArchConfig,
|
TextEncoderArchConfig,
|
||||||
TextEncoderConfig,
|
TextEncoderConfig,
|
||||||
)
|
)
|
||||||
|
from sglang.multimodal_gen.utils import StoreBoolean
|
||||||
|
|
||||||
|
|
||||||
def _is_transformer_layer(n: str, m) -> bool:
|
def _is_transformer_layer(n: str, m) -> bool:
|
||||||
@@ -84,3 +86,28 @@ class T5Config(TextEncoderConfig):
|
|||||||
arch_config: TextEncoderArchConfig = field(default_factory=T5ArchConfig)
|
arch_config: TextEncoderArchConfig = field(default_factory=T5ArchConfig)
|
||||||
|
|
||||||
prefix: str = "t5"
|
prefix: str = "t5"
|
||||||
|
# Use the SP Group of the transformer as the TP Group of T5.
|
||||||
|
parallel_folding: bool = False
|
||||||
|
# "sp" or "ulysses" or "ring"
|
||||||
|
parallel_folding_mode: str = "sp"
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def add_cli_args(
|
||||||
|
parser: argparse.ArgumentParser, prefix: str = "t5-config"
|
||||||
|
) -> argparse.ArgumentParser:
|
||||||
|
parser.add_argument(
|
||||||
|
f"--{prefix}.parallel-folding",
|
||||||
|
action=StoreBoolean,
|
||||||
|
dest=f"{prefix.replace('-', '_')}.parallel_folding",
|
||||||
|
default=T5Config.parallel_folding,
|
||||||
|
help="Whether to use parallel folding for T5",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
f"--{prefix}.parallel-folding-mode",
|
||||||
|
type=str,
|
||||||
|
choices=["sp", "ulysses", "ring"],
|
||||||
|
dest=f"{prefix.replace('-', '_')}.parallel_folding_mode",
|
||||||
|
default=T5Config.parallel_folding_mode,
|
||||||
|
help="Parallel folding mode for T5",
|
||||||
|
)
|
||||||
|
return parser
|
||||||
|
|||||||
@@ -20,12 +20,15 @@ from sglang.multimodal_gen.configs.models import (
|
|||||||
VAEConfig,
|
VAEConfig,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.configs.models.encoders import BaseEncoderOutput
|
from sglang.multimodal_gen.configs.models.encoders import BaseEncoderOutput
|
||||||
|
from sglang.multimodal_gen.configs.models.encoders.t5 import T5Config
|
||||||
from sglang.multimodal_gen.configs.sample.sampling_params import DataType
|
from sglang.multimodal_gen.configs.sample.sampling_params import DataType
|
||||||
from sglang.multimodal_gen.configs.utils import update_config_from_args
|
from sglang.multimodal_gen.configs.utils import update_config_from_args
|
||||||
from sglang.multimodal_gen.runtime.distributed import (
|
from sglang.multimodal_gen.runtime.distributed.communication_op import (
|
||||||
|
sequence_model_parallel_all_gather,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.distributed.parallel_state import (
|
||||||
get_sp_parallel_rank,
|
get_sp_parallel_rank,
|
||||||
get_sp_world_size,
|
get_sp_world_size,
|
||||||
sequence_model_parallel_all_gather,
|
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.models.vision_utils import get_default_height_width
|
from sglang.multimodal_gen.runtime.models.vision_utils import get_default_height_width
|
||||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||||
@@ -492,6 +495,11 @@ class PipelineConfig:
|
|||||||
|
|
||||||
DiTConfig.add_cli_args(parser, prefix=f"{prefix_with_dot}dit-config")
|
DiTConfig.add_cli_args(parser, prefix=f"{prefix_with_dot}dit-config")
|
||||||
|
|
||||||
|
# Add T5 configuration arguments
|
||||||
|
from sglang.multimodal_gen.configs.models.encoders.t5 import T5Config
|
||||||
|
|
||||||
|
T5Config.add_cli_args(parser, prefix=f"{prefix_with_dot}t5-config")
|
||||||
|
|
||||||
return parser
|
return parser
|
||||||
|
|
||||||
def update_config_from_dict(self, args: dict[str, Any], prefix: str = "") -> None:
|
def update_config_from_dict(self, args: dict[str, Any], prefix: str = "") -> None:
|
||||||
@@ -503,6 +511,14 @@ class PipelineConfig:
|
|||||||
update_config_from_args(
|
update_config_from_args(
|
||||||
self.dit_config, args, f"{prefix_with_dot}dit_config", pop_args=True
|
self.dit_config, args, f"{prefix_with_dot}dit_config", pop_args=True
|
||||||
)
|
)
|
||||||
|
for text_encoder_config in self.text_encoder_configs:
|
||||||
|
if isinstance(text_encoder_config, T5Config):
|
||||||
|
update_config_from_args(
|
||||||
|
text_encoder_config,
|
||||||
|
args,
|
||||||
|
f"{prefix_with_dot}t5_config",
|
||||||
|
pop_args=True,
|
||||||
|
)
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def from_kwargs(
|
def from_kwargs(
|
||||||
|
|||||||
@@ -1,7 +1,7 @@
|
|||||||
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
|
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
|
||||||
|
from functools import lru_cache
|
||||||
|
|
||||||
# SPDX-License-Identifier: Apache-2.0
|
from sglang.multimodal_gen.configs.models.encoders import TextEncoderConfig
|
||||||
|
|
||||||
from sglang.multimodal_gen.runtime.distributed.communication_op import *
|
from sglang.multimodal_gen.runtime.distributed.communication_op import *
|
||||||
from sglang.multimodal_gen.runtime.distributed.group_coordinator import (
|
from sglang.multimodal_gen.runtime.distributed.group_coordinator import (
|
||||||
get_local_torch_device,
|
get_local_torch_device,
|
||||||
@@ -27,6 +27,9 @@ from sglang.multimodal_gen.runtime.distributed.parallel_state import (
|
|||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.distributed.utils import *
|
from sglang.multimodal_gen.runtime.distributed.utils import *
|
||||||
|
|
||||||
|
# SPDX-License-Identifier: Apache-2.0
|
||||||
|
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
# Initialization
|
# Initialization
|
||||||
"init_distributed_environment",
|
"init_distributed_environment",
|
||||||
@@ -53,3 +56,16 @@ __all__ = [
|
|||||||
# Get torch device
|
# Get torch device
|
||||||
"get_local_torch_device",
|
"get_local_torch_device",
|
||||||
]
|
]
|
||||||
|
|
||||||
|
|
||||||
|
def _get_folding_tp_group(
|
||||||
|
config: TextEncoderConfig,
|
||||||
|
) -> torch.distributed.ProcessGroup | None:
|
||||||
|
if config.parallel_folding:
|
||||||
|
if config.parallel_folding_mode == "sp":
|
||||||
|
return get_sp_group()
|
||||||
|
elif config.parallel_folding_mode == "ulysses":
|
||||||
|
return get_sp_group().ulysses_group
|
||||||
|
elif config.parallel_folding_mode == "ring":
|
||||||
|
return get_sp_group().ring_group
|
||||||
|
return get_tp_group()
|
||||||
|
|||||||
@@ -4,7 +4,7 @@
|
|||||||
# Adapted from https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/distributed/communication_op.py
|
# Adapted from https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/distributed/communication_op.py
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
import torch.distributed
|
import torch.distributed as dist
|
||||||
|
|
||||||
from sglang.multimodal_gen.runtime.distributed.parallel_state import (
|
from sglang.multimodal_gen.runtime.distributed.parallel_state import (
|
||||||
get_cfg_group,
|
get_cfg_group,
|
||||||
@@ -13,16 +13,20 @@ from sglang.multimodal_gen.runtime.distributed.parallel_state import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def tensor_model_parallel_all_reduce(input_: torch.Tensor) -> torch.Tensor:
|
def tensor_model_parallel_all_reduce(
|
||||||
|
input_: torch.Tensor, tp_group: dist.ProcessGroup = None
|
||||||
|
) -> torch.Tensor:
|
||||||
"""All-reduce the input tensor across model parallel group."""
|
"""All-reduce the input tensor across model parallel group."""
|
||||||
return get_tp_group().all_reduce(input_)
|
tp_group = tp_group or get_tp_group()
|
||||||
|
return tp_group.all_reduce(input_)
|
||||||
|
|
||||||
|
|
||||||
def tensor_model_parallel_all_gather(
|
def tensor_model_parallel_all_gather(
|
||||||
input_: torch.Tensor, dim: int = -1
|
input_: torch.Tensor, dim: int = -1, tp_group: dist.ProcessGroup = None
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
"""All-gather the input tensor across model parallel group."""
|
"""All-gather the input tensor across model parallel group."""
|
||||||
return get_tp_group().all_gather(input_, dim)
|
tp_group = tp_group or get_tp_group()
|
||||||
|
return tp_group.all_gather(input_, dim)
|
||||||
|
|
||||||
|
|
||||||
# TODO: remove model, make it sequence_parallel
|
# TODO: remove model, make it sequence_parallel
|
||||||
|
|||||||
@@ -6,13 +6,13 @@
|
|||||||
from abc import abstractmethod
|
from abc import abstractmethod
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
import torch.distributed as dist
|
||||||
import torch.nn.functional as F
|
import torch.nn.functional as F
|
||||||
from torch.nn.parameter import Parameter
|
from torch.nn.parameter import Parameter
|
||||||
|
|
||||||
from sglang.multimodal_gen.runtime.distributed import (
|
from sglang.multimodal_gen.runtime.distributed import (
|
||||||
divide,
|
divide,
|
||||||
get_tp_rank,
|
get_tp_group,
|
||||||
get_tp_world_size,
|
|
||||||
split_tensor_along_last_dim,
|
split_tensor_along_last_dim,
|
||||||
tensor_model_parallel_all_gather,
|
tensor_model_parallel_all_gather,
|
||||||
tensor_model_parallel_all_reduce,
|
tensor_model_parallel_all_reduce,
|
||||||
@@ -21,6 +21,7 @@ from sglang.multimodal_gen.runtime.layers.quantization.base_config import (
|
|||||||
QuantizationConfig,
|
QuantizationConfig,
|
||||||
QuantizeMethodBase,
|
QuantizeMethodBase,
|
||||||
)
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.layers.utils import get_group_rank, get_group_size
|
||||||
|
|
||||||
# yapf: disable
|
# yapf: disable
|
||||||
from sglang.multimodal_gen.runtime.models.parameter import (
|
from sglang.multimodal_gen.runtime.models.parameter import (
|
||||||
@@ -321,9 +322,12 @@ class ColumnParallelLinear(LinearBase):
|
|||||||
quant_config: QuantizationConfig | None = None,
|
quant_config: QuantizationConfig | None = None,
|
||||||
output_sizes: list[int] | None = None,
|
output_sizes: list[int] | None = None,
|
||||||
prefix: str = "",
|
prefix: str = "",
|
||||||
|
tp_group: dist.ProcessGroup = None,
|
||||||
):
|
):
|
||||||
# Divide the weight matrix along the last dimension.
|
# Divide the weight matrix along the last dimension.
|
||||||
self.tp_size = get_tp_world_size()
|
self.tp_group = tp_group or get_tp_group()
|
||||||
|
self.tp_size = get_group_size(self.tp_group)
|
||||||
|
self.tp_rank = get_group_rank(self.tp_group)
|
||||||
self.input_size_per_partition = input_size
|
self.input_size_per_partition = input_size
|
||||||
self.output_size_per_partition = divide(output_size, self.tp_size)
|
self.output_size_per_partition = divide(output_size, self.tp_size)
|
||||||
self.output_partition_sizes = [self.output_size_per_partition]
|
self.output_partition_sizes = [self.output_size_per_partition]
|
||||||
@@ -374,7 +378,7 @@ class ColumnParallelLinear(LinearBase):
|
|||||||
self.register_parameter("bias", None)
|
self.register_parameter("bias", None)
|
||||||
|
|
||||||
def weight_loader(self, param: Parameter, loaded_weight: torch.Tensor) -> None:
|
def weight_loader(self, param: Parameter, loaded_weight: torch.Tensor) -> None:
|
||||||
tp_rank = get_tp_rank()
|
tp_rank = self.tp_rank
|
||||||
output_dim = getattr(param, "output_dim", None)
|
output_dim = getattr(param, "output_dim", None)
|
||||||
|
|
||||||
is_sharded_weight = getattr(param, "is_sharded_weight", False)
|
is_sharded_weight = getattr(param, "is_sharded_weight", False)
|
||||||
@@ -410,7 +414,9 @@ class ColumnParallelLinear(LinearBase):
|
|||||||
output_parallel = self.quant_method.apply(self, input_, bias)
|
output_parallel = self.quant_method.apply(self, input_, bias)
|
||||||
if self.gather_output:
|
if self.gather_output:
|
||||||
# All-gather across the partitions.
|
# All-gather across the partitions.
|
||||||
output = tensor_model_parallel_all_gather(output_parallel)
|
output = tensor_model_parallel_all_gather(
|
||||||
|
output_parallel, tp_group=self.tp_group
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
output = output_parallel
|
output = output_parallel
|
||||||
output_bias = self.bias if self.skip_bias_add else None
|
output_bias = self.bias if self.skip_bias_add else None
|
||||||
@@ -420,7 +426,7 @@ class ColumnParallelLinear(LinearBase):
|
|||||||
s = f"in_features={self.input_size}"
|
s = f"in_features={self.input_size}"
|
||||||
s += f", output_features={self.output_size_per_partition}"
|
s += f", output_features={self.output_size_per_partition}"
|
||||||
s += f", bias={self.bias is not None}"
|
s += f", bias={self.bias is not None}"
|
||||||
s += f", tp_size={get_tp_world_size()}"
|
s += f", tp_size={self.tp_size}"
|
||||||
s += f", gather_output={self.gather_output}"
|
s += f", gather_output={self.gather_output}"
|
||||||
return s
|
return s
|
||||||
|
|
||||||
@@ -458,10 +464,8 @@ class MergedColumnParallelLinear(ColumnParallelLinear):
|
|||||||
params_dtype: torch.dtype | None = None,
|
params_dtype: torch.dtype | None = None,
|
||||||
quant_config: QuantizationConfig | None = None,
|
quant_config: QuantizationConfig | None = None,
|
||||||
prefix: str = "",
|
prefix: str = "",
|
||||||
|
tp_group: dist.ProcessGroup = None,
|
||||||
):
|
):
|
||||||
self.output_sizes = output_sizes
|
|
||||||
tp_size = get_tp_world_size()
|
|
||||||
assert all(output_size % tp_size == 0 for output_size in output_sizes)
|
|
||||||
super().__init__(
|
super().__init__(
|
||||||
input_size=input_size,
|
input_size=input_size,
|
||||||
output_size=sum(output_sizes),
|
output_size=sum(output_sizes),
|
||||||
@@ -471,7 +475,10 @@ class MergedColumnParallelLinear(ColumnParallelLinear):
|
|||||||
params_dtype=params_dtype,
|
params_dtype=params_dtype,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
prefix=prefix,
|
prefix=prefix,
|
||||||
|
tp_group=tp_group,
|
||||||
)
|
)
|
||||||
|
self.output_sizes = output_sizes
|
||||||
|
assert all(output_size % self.tp_size == 0 for output_size in output_sizes)
|
||||||
|
|
||||||
def weight_loader(
|
def weight_loader(
|
||||||
self,
|
self,
|
||||||
@@ -512,8 +519,8 @@ class MergedColumnParallelLinear(ColumnParallelLinear):
|
|||||||
return
|
return
|
||||||
|
|
||||||
assert loaded_shard_id < len(self.output_sizes)
|
assert loaded_shard_id < len(self.output_sizes)
|
||||||
tp_rank = get_tp_rank()
|
tp_rank = self.tp_rank
|
||||||
tp_size = get_tp_world_size()
|
tp_size = self.tp_size
|
||||||
if output_dim is not None:
|
if output_dim is not None:
|
||||||
shard_offset = sum(self.output_sizes[:loaded_shard_id]) // tp_size
|
shard_offset = sum(self.output_sizes[:loaded_shard_id]) // tp_size
|
||||||
shard_size = self.output_sizes[loaded_shard_id] // tp_size
|
shard_size = self.output_sizes[loaded_shard_id] // tp_size
|
||||||
@@ -607,7 +614,7 @@ class MergedColumnParallelLinear(ColumnParallelLinear):
|
|||||||
|
|
||||||
assert loaded_shard_id < len(self.output_sizes)
|
assert loaded_shard_id < len(self.output_sizes)
|
||||||
|
|
||||||
tp_size = get_tp_world_size()
|
tp_size = self.tp_size
|
||||||
|
|
||||||
if isinstance(param, BlockQuantScaleParameter):
|
if isinstance(param, BlockQuantScaleParameter):
|
||||||
raise NotImplementedError("FP8 is not implemented yet")
|
raise NotImplementedError("FP8 is not implemented yet")
|
||||||
@@ -674,6 +681,7 @@ class QKVParallelLinear(ColumnParallelLinear):
|
|||||||
params_dtype: torch.dtype | None = None,
|
params_dtype: torch.dtype | None = None,
|
||||||
quant_config: QuantizationConfig | None = None,
|
quant_config: QuantizationConfig | None = None,
|
||||||
prefix: str = "",
|
prefix: str = "",
|
||||||
|
tp_group: dist.ProcessGroup = None,
|
||||||
):
|
):
|
||||||
self.hidden_size = hidden_size
|
self.hidden_size = hidden_size
|
||||||
self.head_size = head_size
|
self.head_size = head_size
|
||||||
@@ -682,7 +690,8 @@ class QKVParallelLinear(ColumnParallelLinear):
|
|||||||
total_num_kv_heads = total_num_heads
|
total_num_kv_heads = total_num_heads
|
||||||
self.total_num_kv_heads = total_num_kv_heads
|
self.total_num_kv_heads = total_num_kv_heads
|
||||||
# Divide the weight matrix along the last dimension.
|
# Divide the weight matrix along the last dimension.
|
||||||
tp_size = get_tp_world_size()
|
tp_group = tp_group or get_tp_group()
|
||||||
|
tp_size = get_group_size(tp_group)
|
||||||
self.num_heads = divide(self.total_num_heads, tp_size)
|
self.num_heads = divide(self.total_num_heads, tp_size)
|
||||||
if tp_size >= self.total_num_kv_heads:
|
if tp_size >= self.total_num_kv_heads:
|
||||||
self.num_kv_heads = 1
|
self.num_kv_heads = 1
|
||||||
@@ -709,6 +718,7 @@ class QKVParallelLinear(ColumnParallelLinear):
|
|||||||
params_dtype=params_dtype,
|
params_dtype=params_dtype,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
prefix=prefix,
|
prefix=prefix,
|
||||||
|
tp_group=tp_group,
|
||||||
)
|
)
|
||||||
|
|
||||||
def _get_shard_offset_mapping(self, loaded_shard_id: str) -> int | None:
|
def _get_shard_offset_mapping(self, loaded_shard_id: str) -> int | None:
|
||||||
@@ -852,7 +862,7 @@ class QKVParallelLinear(ColumnParallelLinear):
|
|||||||
self.weight_loader(param, loaded_weight_shard, shard_id)
|
self.weight_loader(param, loaded_weight_shard, shard_id)
|
||||||
return
|
return
|
||||||
|
|
||||||
tp_rank = get_tp_rank()
|
tp_rank = self.tp_rank
|
||||||
assert loaded_shard_id in ["q", "k", "v"]
|
assert loaded_shard_id in ["q", "k", "v"]
|
||||||
|
|
||||||
# If output dim is defined, use the default loading process.
|
# If output dim is defined, use the default loading process.
|
||||||
@@ -944,10 +954,12 @@ class RowParallelLinear(LinearBase):
|
|||||||
reduce_results: bool = True,
|
reduce_results: bool = True,
|
||||||
quant_config: QuantizationConfig | None = None,
|
quant_config: QuantizationConfig | None = None,
|
||||||
prefix: str = "",
|
prefix: str = "",
|
||||||
|
tp_group: dist.ProcessGroup = None,
|
||||||
):
|
):
|
||||||
# Divide the weight matrix along the first dimension.
|
# Divide the weight matrix along the first dimension.
|
||||||
self.tp_rank = get_tp_rank()
|
self.tp_group = tp_group or get_tp_group()
|
||||||
self.tp_size = get_tp_world_size()
|
self.tp_rank = get_group_rank(self.tp_group)
|
||||||
|
self.tp_size = get_group_size(self.tp_group)
|
||||||
self.input_size_per_partition = divide(input_size, self.tp_size)
|
self.input_size_per_partition = divide(input_size, self.tp_size)
|
||||||
self.output_size_per_partition = output_size
|
self.output_size_per_partition = output_size
|
||||||
self.output_partition_sizes = [output_size]
|
self.output_partition_sizes = [output_size]
|
||||||
@@ -992,7 +1004,7 @@ class RowParallelLinear(LinearBase):
|
|||||||
self.register_parameter("bias", None)
|
self.register_parameter("bias", None)
|
||||||
|
|
||||||
def weight_loader(self, param: Parameter, loaded_weight: torch.Tensor):
|
def weight_loader(self, param: Parameter, loaded_weight: torch.Tensor):
|
||||||
tp_rank = get_tp_rank()
|
tp_rank = self.tp_rank
|
||||||
input_dim = getattr(param, "input_dim", None)
|
input_dim = getattr(param, "input_dim", None)
|
||||||
is_sharded_weight = getattr(param, "is_sharded_weight", False)
|
is_sharded_weight = getattr(param, "is_sharded_weight", False)
|
||||||
# bitsandbytes loads the weights of the specific portion
|
# bitsandbytes loads the weights of the specific portion
|
||||||
@@ -1027,7 +1039,7 @@ class RowParallelLinear(LinearBase):
|
|||||||
if self.input_is_parallel:
|
if self.input_is_parallel:
|
||||||
input_parallel = input_
|
input_parallel = input_
|
||||||
else:
|
else:
|
||||||
tp_rank = get_tp_rank()
|
tp_rank = self.tp_rank
|
||||||
splitted_input = split_tensor_along_last_dim(
|
splitted_input = split_tensor_along_last_dim(
|
||||||
input_, num_partitions=self.tp_size
|
input_, num_partitions=self.tp_size
|
||||||
)
|
)
|
||||||
@@ -1040,7 +1052,9 @@ class RowParallelLinear(LinearBase):
|
|||||||
bias_ = None if (self.tp_rank > 0 or self.skip_bias_add) else self.bias
|
bias_ = None if (self.tp_rank > 0 or self.skip_bias_add) else self.bias
|
||||||
output_parallel = self.quant_method.apply(self, input_parallel, bias=bias_)
|
output_parallel = self.quant_method.apply(self, input_parallel, bias=bias_)
|
||||||
if self.reduce_results and self.tp_size > 1:
|
if self.reduce_results and self.tp_size > 1:
|
||||||
output = tensor_model_parallel_all_reduce(output_parallel)
|
output = tensor_model_parallel_all_reduce(
|
||||||
|
output_parallel, tp_group=self.tp_group
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
output = output_parallel
|
output = output_parallel
|
||||||
|
|
||||||
|
|||||||
@@ -12,6 +12,24 @@ from torch.library import Library
|
|||||||
from sglang.multimodal_gen.runtime.platforms import current_platform
|
from sglang.multimodal_gen.runtime.platforms import current_platform
|
||||||
|
|
||||||
|
|
||||||
|
def get_group_size(group) -> int:
|
||||||
|
if hasattr(group, "world_size"):
|
||||||
|
return group.world_size # GroupCoordinator
|
||||||
|
elif hasattr(group, "size") and callable(getattr(group, "size", None)):
|
||||||
|
return group.size() # ProcessGroup
|
||||||
|
else:
|
||||||
|
raise ValueError(f"Unsupported group type: {type(group)}")
|
||||||
|
|
||||||
|
|
||||||
|
def get_group_rank(group) -> int:
|
||||||
|
if hasattr(group, "rank_in_group"):
|
||||||
|
return group.rank_in_group # GroupCoordinator
|
||||||
|
elif hasattr(group, "rank") and callable(getattr(group, "rank", None)):
|
||||||
|
return group.rank() # ProcessGroup
|
||||||
|
else:
|
||||||
|
raise ValueError(f"Unsupported group type: {type(group)}")
|
||||||
|
|
||||||
|
|
||||||
def get_token_bin_counts_and_mask(
|
def get_token_bin_counts_and_mask(
|
||||||
tokens: torch.Tensor,
|
tokens: torch.Tensor,
|
||||||
vocab_size: int,
|
vocab_size: int,
|
||||||
|
|||||||
@@ -6,13 +6,13 @@ from collections.abc import Sequence
|
|||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
import torch.distributed as dist
|
||||||
import torch.nn.functional as F
|
import torch.nn.functional as F
|
||||||
from torch.nn.parameter import Parameter, UninitializedParameter
|
from torch.nn.parameter import Parameter, UninitializedParameter
|
||||||
|
|
||||||
from sglang.multimodal_gen.runtime.distributed import (
|
from sglang.multimodal_gen.runtime.distributed import (
|
||||||
divide,
|
divide,
|
||||||
get_tp_rank,
|
get_tp_group,
|
||||||
get_tp_world_size,
|
|
||||||
tensor_model_parallel_all_reduce,
|
tensor_model_parallel_all_reduce,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.layers.quantization.base_config import (
|
from sglang.multimodal_gen.runtime.layers.quantization.base_config import (
|
||||||
@@ -20,6 +20,7 @@ from sglang.multimodal_gen.runtime.layers.quantization.base_config import (
|
|||||||
QuantizeMethodBase,
|
QuantizeMethodBase,
|
||||||
method_has_implemented_embedding,
|
method_has_implemented_embedding,
|
||||||
)
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.layers.utils import get_group_rank, get_group_size
|
||||||
from sglang.multimodal_gen.runtime.models.parameter import BasevLLMParameter
|
from sglang.multimodal_gen.runtime.models.parameter import BasevLLMParameter
|
||||||
from sglang.multimodal_gen.runtime.models.utils import set_weight_attrs
|
from sglang.multimodal_gen.runtime.models.utils import set_weight_attrs
|
||||||
from sglang.multimodal_gen.runtime.platforms import current_platform
|
from sglang.multimodal_gen.runtime.platforms import current_platform
|
||||||
@@ -220,12 +221,15 @@ class VocabParallelEmbedding(torch.nn.Module):
|
|||||||
padding_size: int = DEFAULT_VOCAB_PADDING_SIZE,
|
padding_size: int = DEFAULT_VOCAB_PADDING_SIZE,
|
||||||
quant_config: QuantizationConfig | None = None,
|
quant_config: QuantizationConfig | None = None,
|
||||||
prefix: str = "",
|
prefix: str = "",
|
||||||
|
tp_group: dist.ProcessGroup = None,
|
||||||
):
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
|
|
||||||
# Keep the input dimensions.
|
# Keep the input dimensions.
|
||||||
tp_rank = get_tp_rank()
|
tp_group = tp_group or get_tp_group()
|
||||||
self.tp_size = get_tp_world_size()
|
tp_rank = get_group_rank(tp_group)
|
||||||
|
self.tp_size = get_group_size(tp_group)
|
||||||
|
self.tp_group = tp_group
|
||||||
self.num_embeddings = num_embeddings
|
self.num_embeddings = num_embeddings
|
||||||
self.padding_size = padding_size
|
self.padding_size = padding_size
|
||||||
self.org_vocab_size = org_num_embeddings or num_embeddings
|
self.org_vocab_size = org_num_embeddings or num_embeddings
|
||||||
@@ -468,7 +472,9 @@ class VocabParallelEmbedding(torch.nn.Module):
|
|||||||
if self.tp_size > 1:
|
if self.tp_size > 1:
|
||||||
output_parallel.masked_fill_(input_mask.unsqueeze(-1), 0)
|
output_parallel.masked_fill_(input_mask.unsqueeze(-1), 0)
|
||||||
# Reduce across all the model parallel GPUs.
|
# Reduce across all the model parallel GPUs.
|
||||||
output = tensor_model_parallel_all_reduce(output_parallel)
|
output = tensor_model_parallel_all_reduce(
|
||||||
|
output_parallel, tp_group=self.tp_group
|
||||||
|
)
|
||||||
return output
|
return output
|
||||||
|
|
||||||
def extra_repr(self) -> str:
|
def extra_repr(self) -> str:
|
||||||
|
|||||||
@@ -739,8 +739,11 @@ class WanTransformer3DModel(CachableDiT, OffloadableDiTMixin):
|
|||||||
elementwise_affine=False,
|
elementwise_affine=False,
|
||||||
dtype=torch.float32,
|
dtype=torch.float32,
|
||||||
)
|
)
|
||||||
self.proj_out = nn.Linear(
|
self.proj_out = ColumnParallelLinear(
|
||||||
inner_dim, config.out_channels * math.prod(config.patch_size)
|
inner_dim,
|
||||||
|
config.out_channels * math.prod(config.patch_size),
|
||||||
|
bias=True,
|
||||||
|
gather_output=True,
|
||||||
)
|
)
|
||||||
self.scale_shift_table = nn.Parameter(
|
self.scale_shift_table = nn.Parameter(
|
||||||
torch.randn(1, 2, inner_dim) / inner_dim**0.5
|
torch.randn(1, 2, inner_dim) / inner_dim**0.5
|
||||||
@@ -891,7 +894,7 @@ class WanTransformer3DModel(CachableDiT, OffloadableDiTMixin):
|
|||||||
shift, scale = (self.scale_shift_table + temb.unsqueeze(1)).chunk(2, dim=1)
|
shift, scale = (self.scale_shift_table + temb.unsqueeze(1)).chunk(2, dim=1)
|
||||||
|
|
||||||
hidden_states = self.norm_out(hidden_states, shift, scale)
|
hidden_states = self.norm_out(hidden_states, shift, scale)
|
||||||
hidden_states = self.proj_out(hidden_states)
|
hidden_states, _ = self.proj_out(hidden_states)
|
||||||
|
|
||||||
hidden_states = hidden_states.reshape(
|
hidden_states = hidden_states.reshape(
|
||||||
batch_size,
|
batch_size,
|
||||||
|
|||||||
@@ -30,7 +30,7 @@ import torch.nn.functional as F
|
|||||||
from torch import nn
|
from torch import nn
|
||||||
|
|
||||||
from sglang.multimodal_gen.configs.models.encoders import BaseEncoderOutput, T5Config
|
from sglang.multimodal_gen.configs.models.encoders import BaseEncoderOutput, T5Config
|
||||||
from sglang.multimodal_gen.runtime.distributed import get_tp_rank, get_tp_world_size
|
from sglang.multimodal_gen.runtime.distributed import _get_folding_tp_group
|
||||||
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.layernorm import RMSNorm
|
from sglang.multimodal_gen.runtime.layers.layernorm import RMSNorm
|
||||||
from sglang.multimodal_gen.runtime.layers.linear import (
|
from sglang.multimodal_gen.runtime.layers.linear import (
|
||||||
@@ -39,6 +39,7 @@ from sglang.multimodal_gen.runtime.layers.linear import (
|
|||||||
RowParallelLinear,
|
RowParallelLinear,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.layers.quantization import QuantizationConfig
|
from sglang.multimodal_gen.runtime.layers.quantization import QuantizationConfig
|
||||||
|
from sglang.multimodal_gen.runtime.layers.utils import get_group_rank, get_group_size
|
||||||
from sglang.multimodal_gen.runtime.layers.vocab_parallel_embedding import (
|
from sglang.multimodal_gen.runtime.layers.vocab_parallel_embedding import (
|
||||||
VocabParallelEmbedding,
|
VocabParallelEmbedding,
|
||||||
)
|
)
|
||||||
@@ -74,9 +75,16 @@ class T5DenseActDense(nn.Module):
|
|||||||
self, config: T5Config, quant_config: QuantizationConfig | None = None
|
self, config: T5Config, quant_config: QuantizationConfig | None = None
|
||||||
):
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.wi = MergedColumnParallelLinear(config.d_model, [config.d_ff], bias=False)
|
tp_group = _get_folding_tp_group(config)
|
||||||
|
self.wi = MergedColumnParallelLinear(
|
||||||
|
config.d_model, [config.d_ff], bias=False, tp_group=tp_group
|
||||||
|
)
|
||||||
self.wo = RowParallelLinear(
|
self.wo = RowParallelLinear(
|
||||||
config.d_ff, config.d_model, bias=False, quant_config=quant_config
|
config.d_ff,
|
||||||
|
config.d_model,
|
||||||
|
bias=False,
|
||||||
|
quant_config=quant_config,
|
||||||
|
tp_group=tp_group,
|
||||||
)
|
)
|
||||||
self.act = get_act_fn(config.dense_act_fn)
|
self.act = get_act_fn(config.dense_act_fn)
|
||||||
|
|
||||||
@@ -93,16 +101,29 @@ class T5DenseGatedActDense(nn.Module):
|
|||||||
self, config: T5Config, quant_config: QuantizationConfig | None = None
|
self, config: T5Config, quant_config: QuantizationConfig | None = None
|
||||||
):
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
|
tp_group = _get_folding_tp_group(config)
|
||||||
self.wi_0 = MergedColumnParallelLinear(
|
self.wi_0 = MergedColumnParallelLinear(
|
||||||
config.d_model, [config.d_ff], bias=False, quant_config=quant_config
|
config.d_model,
|
||||||
|
[config.d_ff],
|
||||||
|
bias=False,
|
||||||
|
quant_config=quant_config,
|
||||||
|
tp_group=tp_group,
|
||||||
)
|
)
|
||||||
self.wi_1 = MergedColumnParallelLinear(
|
self.wi_1 = MergedColumnParallelLinear(
|
||||||
config.d_model, [config.d_ff], bias=False, quant_config=quant_config
|
config.d_model,
|
||||||
|
[config.d_ff],
|
||||||
|
bias=False,
|
||||||
|
quant_config=quant_config,
|
||||||
|
tp_group=tp_group,
|
||||||
)
|
)
|
||||||
# Should not run in fp16 unless mixed-precision is used,
|
# Should not run in fp16 unless mixed-precision is used,
|
||||||
# see https://github.com/huggingface/transformers/issues/20287.
|
# see https://github.com/huggingface/transformers/issues/20287.
|
||||||
self.wo = RowParallelLinear(
|
self.wo = RowParallelLinear(
|
||||||
config.d_ff, config.d_model, bias=False, quant_config=quant_config
|
config.d_ff,
|
||||||
|
config.d_model,
|
||||||
|
bias=False,
|
||||||
|
quant_config=quant_config,
|
||||||
|
tp_group=tp_group,
|
||||||
)
|
)
|
||||||
self.act = get_act_fn(config.dense_act_fn)
|
self.act = get_act_fn(config.dense_act_fn)
|
||||||
|
|
||||||
@@ -176,9 +197,10 @@ class T5Attention(nn.Module):
|
|||||||
self.total_num_heads = self.total_num_kv_heads = config.num_heads
|
self.total_num_heads = self.total_num_kv_heads = config.num_heads
|
||||||
|
|
||||||
# Partition heads across multiple tensor parallel GPUs.
|
# Partition heads across multiple tensor parallel GPUs.
|
||||||
tp_world_size = get_tp_world_size()
|
self.tp_group = _get_folding_tp_group(config)
|
||||||
assert config.num_heads % tp_world_size == 0
|
self.tp_world_size = get_group_size(self.tp_group)
|
||||||
self.n_heads = config.num_heads // tp_world_size
|
assert config.num_heads % self.tp_world_size == 0
|
||||||
|
self.n_heads = config.num_heads // self.tp_world_size
|
||||||
|
|
||||||
self.inner_dim = self.n_heads * self.key_value_proj_dim
|
self.inner_dim = self.n_heads * self.key_value_proj_dim
|
||||||
# No GQA in t5.
|
# No GQA in t5.
|
||||||
@@ -192,6 +214,7 @@ class T5Attention(nn.Module):
|
|||||||
bias=False,
|
bias=False,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
prefix=f"{prefix}.qkv_proj",
|
prefix=f"{prefix}.qkv_proj",
|
||||||
|
tp_group=self.tp_group,
|
||||||
)
|
)
|
||||||
|
|
||||||
self.attn = T5MultiHeadAttention()
|
self.attn = T5MultiHeadAttention()
|
||||||
@@ -203,6 +226,7 @@ class T5Attention(nn.Module):
|
|||||||
org_num_embeddings=self.relative_attention_num_buckets,
|
org_num_embeddings=self.relative_attention_num_buckets,
|
||||||
padding_size=self.relative_attention_num_buckets,
|
padding_size=self.relative_attention_num_buckets,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
|
tp_group=self.tp_group,
|
||||||
)
|
)
|
||||||
self.o = RowParallelLinear(
|
self.o = RowParallelLinear(
|
||||||
self.total_num_heads * self.key_value_proj_dim,
|
self.total_num_heads * self.key_value_proj_dim,
|
||||||
@@ -210,6 +234,7 @@ class T5Attention(nn.Module):
|
|||||||
bias=False,
|
bias=False,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
prefix=f"{prefix}.o_proj",
|
prefix=f"{prefix}.o_proj",
|
||||||
|
tp_group=self.tp_group,
|
||||||
)
|
)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
@@ -339,8 +364,8 @@ class T5Attention(nn.Module):
|
|||||||
mask_val = -1e4 if current_platform.is_mps() else torch.finfo(q.dtype).min
|
mask_val = -1e4 if current_platform.is_mps() else torch.finfo(q.dtype).min
|
||||||
attn_bias.masked_fill_(attention_mask == 0, mask_val)
|
attn_bias.masked_fill_(attention_mask == 0, mask_val)
|
||||||
|
|
||||||
if get_tp_world_size() > 1:
|
if self.tp_world_size > 1:
|
||||||
rank = get_tp_rank()
|
rank = get_group_rank(self.tp_group)
|
||||||
attn_bias = attn_bias[
|
attn_bias = attn_bias[
|
||||||
:, rank * self.n_heads : (rank + 1) * self.n_heads, :, :
|
:, rank * self.n_heads : (rank + 1) * self.n_heads, :, :
|
||||||
]
|
]
|
||||||
@@ -546,9 +571,12 @@ class T5EncoderModel(TextEncoder):
|
|||||||
super().__init__(config)
|
super().__init__(config)
|
||||||
|
|
||||||
quant_config = None
|
quant_config = None
|
||||||
|
tp_group = _get_folding_tp_group(config)
|
||||||
self.shared = VocabParallelEmbedding(
|
self.shared = VocabParallelEmbedding(
|
||||||
config.vocab_size, config.d_model, org_num_embeddings=config.vocab_size
|
config.vocab_size,
|
||||||
|
config.d_model,
|
||||||
|
org_num_embeddings=config.vocab_size,
|
||||||
|
tp_group=tp_group,
|
||||||
)
|
)
|
||||||
|
|
||||||
self.encoder = T5Stack(
|
self.encoder = T5Stack(
|
||||||
@@ -632,9 +660,12 @@ class UMT5EncoderModel(TextEncoder):
|
|||||||
super().__init__(config)
|
super().__init__(config)
|
||||||
|
|
||||||
quant_config = None
|
quant_config = None
|
||||||
|
tp_group = _get_folding_tp_group(config)
|
||||||
self.shared = VocabParallelEmbedding(
|
self.shared = VocabParallelEmbedding(
|
||||||
config.vocab_size, config.d_model, org_num_embeddings=config.vocab_size
|
config.vocab_size,
|
||||||
|
config.d_model,
|
||||||
|
org_num_embeddings=config.vocab_size,
|
||||||
|
tp_group=tp_group,
|
||||||
)
|
)
|
||||||
|
|
||||||
self.encoder = T5Stack(
|
self.encoder = T5Stack(
|
||||||
|
|||||||
Reference in New Issue
Block a user