[diffusion] support tp for ideogram4 (#27393)
This commit is contained in:
@@ -14,7 +14,7 @@ class LTXVocoderArchConfig(VocoderArchConfig):
|
|||||||
in_channels: int = 128
|
in_channels: int = 128
|
||||||
hidden_channels: int = 1024
|
hidden_channels: int = 1024
|
||||||
out_channels: int = 2
|
out_channels: int = 2
|
||||||
upsample_kernel_sizes: List[int] = field(default_factory=lambda: [3, 7, 11])
|
upsample_kernel_sizes: List[int] = field(default_factory=lambda: [16, 15, 8, 4, 4])
|
||||||
upsample_factors: List[int] = field(default_factory=lambda: [6, 5, 2, 2, 2])
|
upsample_factors: List[int] = field(default_factory=lambda: [6, 5, 2, 2, 2])
|
||||||
resnet_kernel_sizes: List[int] = field(default_factory=lambda: [3, 7, 11])
|
resnet_kernel_sizes: List[int] = field(default_factory=lambda: [3, 7, 11])
|
||||||
resnet_dilations: List[List[int]] = field(
|
resnet_dilations: List[List[int]] = field(
|
||||||
|
|||||||
@@ -161,9 +161,9 @@ class BitsAndBytesLinearMethod(LinearMethodBase):
|
|||||||
params_dtype: torch.dtype,
|
params_dtype: torch.dtype,
|
||||||
**extra_weight_attrs,
|
**extra_weight_attrs,
|
||||||
) -> None:
|
) -> None:
|
||||||
del input_size, output_size
|
|
||||||
quant_ratio = _calculate_quant_ratio(params_dtype)
|
quant_ratio = _calculate_quant_ratio(params_dtype)
|
||||||
total_size = input_size_per_partition * sum(output_partition_sizes)
|
output_size_per_partition = sum(output_partition_sizes)
|
||||||
|
total_size = input_size_per_partition * output_size_per_partition
|
||||||
if total_size % quant_ratio != 0:
|
if total_size % quant_ratio != 0:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
"The input size is not aligned with the quantized weight shape."
|
"The input size is not aligned with the quantized weight shape."
|
||||||
@@ -180,6 +180,18 @@ class BitsAndBytesLinearMethod(LinearMethodBase):
|
|||||||
"output_dim": 0,
|
"output_dim": 0,
|
||||||
"pack_factor": quant_ratio,
|
"pack_factor": quant_ratio,
|
||||||
"use_bitsandbytes_4bit": True,
|
"use_bitsandbytes_4bit": True,
|
||||||
|
"bnb_full_shape": (output_size, input_size),
|
||||||
|
"bnb_local_shape": (
|
||||||
|
output_size_per_partition,
|
||||||
|
input_size_per_partition,
|
||||||
|
),
|
||||||
|
"bnb_output_shard_start": getattr(layer, "tp_rank", 0)
|
||||||
|
* output_size_per_partition,
|
||||||
|
"bnb_input_shard_start": (
|
||||||
|
0
|
||||||
|
if input_size_per_partition == input_size
|
||||||
|
else getattr(layer, "tp_rank", 0) * input_size_per_partition
|
||||||
|
),
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
layer.register_parameter("weight", qweight)
|
layer.register_parameter("weight", qweight)
|
||||||
@@ -377,7 +389,49 @@ def attach_bitsandbytes_4bit_quant_states(
|
|||||||
if param is None:
|
if param is None:
|
||||||
raise ValueError(f"Parameter {param_name} not found in the model.")
|
raise ValueError(f"Parameter {param_name} not found in the model.")
|
||||||
|
|
||||||
|
quant_state = _maybe_shard_bitsandbytes_4bit_quant_state(param, quant_state)
|
||||||
state_by_shard = {0: quant_state}
|
state_by_shard = {0: quant_state}
|
||||||
set_weight_attrs(param, {"bnb_quant_state": state_by_shard})
|
set_weight_attrs(param, {"bnb_quant_state": state_by_shard})
|
||||||
offsets = torch.tensor([0, param.numel()]).cpu()
|
offsets = torch.tensor([0, param.numel()]).cpu()
|
||||||
set_weight_attrs(param, {"bnb_shard_offsets": offsets})
|
set_weight_attrs(param, {"bnb_shard_offsets": offsets})
|
||||||
|
|
||||||
|
|
||||||
|
def _maybe_shard_bitsandbytes_4bit_quant_state(
|
||||||
|
param: torch.nn.Parameter,
|
||||||
|
quant_state: Any,
|
||||||
|
) -> Any:
|
||||||
|
full_shape = tuple(getattr(param, "bnb_full_shape", tuple(quant_state.shape or ())))
|
||||||
|
local_shape = tuple(getattr(param, "bnb_local_shape", full_shape))
|
||||||
|
if not full_shape or local_shape == full_shape:
|
||||||
|
return quant_state
|
||||||
|
|
||||||
|
output_start = getattr(param, "bnb_output_shard_start", 0)
|
||||||
|
input_start = getattr(param, "bnb_input_shard_start", 0)
|
||||||
|
if input_start != 0 or local_shape[1] != full_shape[1]:
|
||||||
|
raise NotImplementedError(
|
||||||
|
"bitsandbytes 4-bit TP only supports column-parallel output shards."
|
||||||
|
)
|
||||||
|
if getattr(quant_state, "nested", False):
|
||||||
|
raise NotImplementedError(
|
||||||
|
"bitsandbytes 4-bit TP does not support nested quant states."
|
||||||
|
)
|
||||||
|
|
||||||
|
blocksize = quant_state.blocksize
|
||||||
|
start_elem = output_start * full_shape[1]
|
||||||
|
local_numel = local_shape[0] * local_shape[1]
|
||||||
|
if start_elem % blocksize != 0 or local_numel % blocksize != 0:
|
||||||
|
raise ValueError(
|
||||||
|
"bitsandbytes 4-bit TP shard is not aligned to quantization blocks."
|
||||||
|
)
|
||||||
|
start_block = start_elem // blocksize
|
||||||
|
num_blocks = local_numel // blocksize
|
||||||
|
return type(quant_state)(
|
||||||
|
absmax=quant_state.absmax.narrow(0, start_block, num_blocks).contiguous(),
|
||||||
|
shape=torch.Size(local_shape),
|
||||||
|
code=quant_state.code,
|
||||||
|
blocksize=quant_state.blocksize,
|
||||||
|
quant_type=quant_state.quant_type,
|
||||||
|
dtype=quant_state.dtype,
|
||||||
|
offset=None,
|
||||||
|
state2=None,
|
||||||
|
)
|
||||||
|
|||||||
@@ -4,6 +4,12 @@ import torch
|
|||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
import torch.nn.functional as F
|
import torch.nn.functional as F
|
||||||
|
|
||||||
|
from sglang.multimodal_gen.runtime.distributed import (
|
||||||
|
divide,
|
||||||
|
get_tp_group,
|
||||||
|
tensor_model_parallel_all_gather,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.layers.utils import get_group_rank, get_group_size
|
||||||
from sglang.multimodal_gen.runtime.models.utils import set_weight_attrs
|
from sglang.multimodal_gen.runtime.models.utils import set_weight_attrs
|
||||||
|
|
||||||
FP8_WEIGHT_DTYPE = torch.float8_e4m3fn
|
FP8_WEIGHT_DTYPE = torch.float8_e4m3fn
|
||||||
@@ -66,6 +72,100 @@ class WeightOnlyFP8Linear(nn.Module):
|
|||||||
return F.linear(x.to(compute_dtype), weight, bias)
|
return F.linear(x.to(compute_dtype), weight, bias)
|
||||||
|
|
||||||
|
|
||||||
|
class WeightOnlyFP8ColumnParallelLinear(nn.Module):
|
||||||
|
"""Column-parallel e4m3 FP8 linear with row-wise dequantization."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
in_features: int,
|
||||||
|
out_features: int,
|
||||||
|
bias: bool = True,
|
||||||
|
compute_dtype: torch.dtype | None = None,
|
||||||
|
gather_output: bool = True,
|
||||||
|
tp_group=None,
|
||||||
|
) -> None:
|
||||||
|
super().__init__()
|
||||||
|
self.in_features = in_features
|
||||||
|
self.out_features = out_features
|
||||||
|
self.compute_dtype = compute_dtype
|
||||||
|
self.gather_output = gather_output
|
||||||
|
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.out_features_per_partition = divide(out_features, self.tp_size)
|
||||||
|
self.weight = nn.Parameter(
|
||||||
|
torch.empty(
|
||||||
|
self.out_features_per_partition,
|
||||||
|
in_features,
|
||||||
|
dtype=FP8_WEIGHT_DTYPE,
|
||||||
|
),
|
||||||
|
requires_grad=False,
|
||||||
|
)
|
||||||
|
set_weight_attrs(
|
||||||
|
self.weight,
|
||||||
|
{
|
||||||
|
"output_dim": 0,
|
||||||
|
"weight_loader": self.weight_loader,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
self.weight_scale = nn.Parameter(
|
||||||
|
torch.empty(self.out_features_per_partition, dtype=torch.float32),
|
||||||
|
requires_grad=False,
|
||||||
|
)
|
||||||
|
set_weight_attrs(
|
||||||
|
self.weight_scale,
|
||||||
|
{
|
||||||
|
"missing_param_init": "error",
|
||||||
|
"output_dim": 0,
|
||||||
|
"weight_loader": self.weight_loader,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
if bias:
|
||||||
|
self.bias = nn.Parameter(
|
||||||
|
torch.empty(
|
||||||
|
self.out_features_per_partition,
|
||||||
|
dtype=compute_dtype or torch.get_default_dtype(),
|
||||||
|
),
|
||||||
|
requires_grad=False,
|
||||||
|
)
|
||||||
|
set_weight_attrs(
|
||||||
|
self.bias,
|
||||||
|
{
|
||||||
|
"output_dim": 0,
|
||||||
|
"weight_loader": self.weight_loader,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
self.register_parameter("bias", None)
|
||||||
|
|
||||||
|
def weight_loader(
|
||||||
|
self, param: torch.nn.Parameter, loaded_weight: torch.Tensor
|
||||||
|
) -> None:
|
||||||
|
output_dim = getattr(param, "output_dim", None)
|
||||||
|
if output_dim is not None:
|
||||||
|
shard_size = param.data.shape[output_dim]
|
||||||
|
loaded_weight = loaded_weight.narrow(
|
||||||
|
output_dim, self.tp_rank * shard_size, shard_size
|
||||||
|
)
|
||||||
|
if len(loaded_weight.shape) == 0:
|
||||||
|
loaded_weight = loaded_weight.reshape(1)
|
||||||
|
assert param.data.shape == loaded_weight.shape
|
||||||
|
param.data.copy_(loaded_weight)
|
||||||
|
|
||||||
|
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
|
compute_dtype = self.compute_dtype or x.dtype
|
||||||
|
weight = dequantize_rowwise_fp8_weight(
|
||||||
|
self.weight, self.weight_scale, compute_dtype
|
||||||
|
)
|
||||||
|
bias = self.bias.to(compute_dtype) if self.bias is not None else None
|
||||||
|
output_parallel = F.linear(x.to(compute_dtype), weight, bias)
|
||||||
|
if self.gather_output:
|
||||||
|
return tensor_model_parallel_all_gather(
|
||||||
|
output_parallel, tp_group=self.tp_group
|
||||||
|
)
|
||||||
|
return output_parallel
|
||||||
|
|
||||||
|
|
||||||
def swap_linears_to_weight_only_fp8(module: nn.Module) -> None:
|
def swap_linears_to_weight_only_fp8(module: nn.Module) -> None:
|
||||||
"""Recursively replace nn.Linear with WeightOnlyFP8Linear.
|
"""Recursively replace nn.Linear with WeightOnlyFP8Linear.
|
||||||
|
|
||||||
|
|||||||
@@ -86,7 +86,10 @@ def _make_param_like(
|
|||||||
try:
|
try:
|
||||||
new_param = cls.__new__(cls, tensor, requires_grad=False)
|
new_param = cls.__new__(cls, tensor, requires_grad=False)
|
||||||
except TypeError:
|
except TypeError:
|
||||||
new_param = cls.__new__(cls, tensor)
|
try:
|
||||||
|
new_param = cls.__new__(cls, tensor)
|
||||||
|
except TypeError:
|
||||||
|
new_param = nn.Parameter(tensor, requires_grad=False)
|
||||||
new_param.__dict__.update(actual_param.__dict__)
|
new_param.__dict__.update(actual_param.__dict__)
|
||||||
new_param.requires_grad = False
|
new_param.requires_grad = False
|
||||||
return new_param
|
return new_param
|
||||||
@@ -583,10 +586,15 @@ def load_model_from_full_model_state_dict(
|
|||||||
if cpu_offload:
|
if cpu_offload:
|
||||||
sharded_tensor = sharded_tensor.to("cpu")
|
sharded_tensor = sharded_tensor.to("cpu")
|
||||||
|
|
||||||
requires_grad = False
|
actual_param = param_dict.get(target_param_name)
|
||||||
sharded_sd[target_param_name] = nn.Parameter(
|
if actual_param is not None:
|
||||||
sharded_tensor, requires_grad=requires_grad
|
sharded_sd[target_param_name] = _make_param_like(
|
||||||
)
|
actual_param, sharded_tensor
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
sharded_sd[target_param_name] = nn.Parameter(
|
||||||
|
sharded_tensor, requires_grad=False
|
||||||
|
)
|
||||||
|
|
||||||
model.reverse_param_names_mapping = reverse_param_names_mapping
|
model.reverse_param_names_mapping = reverse_param_names_mapping
|
||||||
|
|
||||||
|
|||||||
@@ -8,15 +8,23 @@ import torch.nn as nn
|
|||||||
import torch.nn.functional as F
|
import torch.nn.functional as F
|
||||||
|
|
||||||
from sglang.multimodal_gen.configs.models.dits.ideogram import Ideogram4DiTConfig
|
from sglang.multimodal_gen.configs.models.dits.ideogram import Ideogram4DiTConfig
|
||||||
|
from sglang.multimodal_gen.runtime.distributed import (
|
||||||
|
get_tp_world_size,
|
||||||
|
model_parallel_is_initialized,
|
||||||
|
)
|
||||||
from sglang.multimodal_gen.runtime.layers.attention import (
|
from sglang.multimodal_gen.runtime.layers.attention import (
|
||||||
USPAttention,
|
USPAttention,
|
||||||
build_varlen_mask_meta,
|
build_varlen_mask_meta,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.layers.linear import ReplicatedLinear
|
from sglang.multimodal_gen.runtime.layers.linear import (
|
||||||
|
ColumnParallelLinear,
|
||||||
|
ReplicatedLinear,
|
||||||
|
)
|
||||||
from sglang.multimodal_gen.runtime.layers.quantization.configs.base_config import (
|
from sglang.multimodal_gen.runtime.layers.quantization.configs.base_config import (
|
||||||
QuantizationConfig,
|
QuantizationConfig,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.layers.quantization.weight_only_fp8 import (
|
from sglang.multimodal_gen.runtime.layers.quantization.weight_only_fp8 import (
|
||||||
|
WeightOnlyFP8ColumnParallelLinear,
|
||||||
WeightOnlyFP8Linear,
|
WeightOnlyFP8Linear,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.layers.rotary_embedding import (
|
from sglang.multimodal_gen.runtime.layers.rotary_embedding import (
|
||||||
@@ -44,6 +52,11 @@ class Ideogram4QuantizedLinear(ReplicatedLinear):
|
|||||||
return super().forward(x)[0]
|
return super().forward(x)[0]
|
||||||
|
|
||||||
|
|
||||||
|
class Ideogram4ColumnParallelLinear(ColumnParallelLinear):
|
||||||
|
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
|
return super().forward(x)[0]
|
||||||
|
|
||||||
|
|
||||||
def _linear(
|
def _linear(
|
||||||
in_features: int,
|
in_features: int,
|
||||||
out_features: int,
|
out_features: int,
|
||||||
@@ -51,8 +64,26 @@ def _linear(
|
|||||||
quant_config: QuantizationConfig | None = None,
|
quant_config: QuantizationConfig | None = None,
|
||||||
prefix: str = "",
|
prefix: str = "",
|
||||||
):
|
):
|
||||||
|
tp_size = get_tp_world_size() if model_parallel_is_initialized() else 1
|
||||||
|
use_column_parallel = tp_size > 1 and out_features % tp_size == 0
|
||||||
if quant_config is None:
|
if quant_config is None:
|
||||||
|
if use_column_parallel:
|
||||||
|
return WeightOnlyFP8ColumnParallelLinear(
|
||||||
|
in_features,
|
||||||
|
out_features,
|
||||||
|
bias=bias,
|
||||||
|
gather_output=True,
|
||||||
|
)
|
||||||
return WeightOnlyFP8Linear(in_features, out_features, bias=bias)
|
return WeightOnlyFP8Linear(in_features, out_features, bias=bias)
|
||||||
|
if use_column_parallel:
|
||||||
|
return Ideogram4ColumnParallelLinear(
|
||||||
|
in_features,
|
||||||
|
out_features,
|
||||||
|
bias=bias,
|
||||||
|
gather_output=True,
|
||||||
|
quant_config=quant_config,
|
||||||
|
prefix=prefix,
|
||||||
|
)
|
||||||
return Ideogram4QuantizedLinear(
|
return Ideogram4QuantizedLinear(
|
||||||
in_features,
|
in_features,
|
||||||
out_features,
|
out_features,
|
||||||
|
|||||||
@@ -11,13 +11,10 @@ from sglang.multimodal_gen.configs.models.encoders.ideogram import (
|
|||||||
Ideogram4TextEncoderConfig,
|
Ideogram4TextEncoderConfig,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.layers.quantization.bitsandbytes import (
|
from sglang.multimodal_gen.runtime.layers.quantization.bitsandbytes import (
|
||||||
|
BitsAndBytesConfig,
|
||||||
attach_bitsandbytes_4bit_quant_states,
|
attach_bitsandbytes_4bit_quant_states,
|
||||||
build_bitsandbytes_4bit_quant_states,
|
build_bitsandbytes_4bit_quant_states,
|
||||||
is_bitsandbytes_4bit_state_name,
|
is_bitsandbytes_4bit_state_name,
|
||||||
swap_linears_to_bitsandbytes_4bit,
|
|
||||||
)
|
|
||||||
from sglang.multimodal_gen.runtime.layers.quantization.weight_only_fp8 import (
|
|
||||||
swap_linears_to_weight_only_fp8,
|
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.loader.weight_utils import default_weight_loader
|
from sglang.multimodal_gen.runtime.loader.weight_utils import default_weight_loader
|
||||||
from sglang.multimodal_gen.runtime.managers.forward_context import set_forward_context
|
from sglang.multimodal_gen.runtime.managers.forward_context import set_forward_context
|
||||||
@@ -36,14 +33,26 @@ class IdeogramQwen3VLTextEncoder(TextEncoder):
|
|||||||
text_config = getattr(arch_config, "text_config")
|
text_config = getattr(arch_config, "text_config")
|
||||||
if isinstance(text_config, dict):
|
if isinstance(text_config, dict):
|
||||||
text_config = Qwen3VLTextConfig(**text_config)
|
text_config = Qwen3VLTextConfig(**text_config)
|
||||||
self.language_model = Qwen3VLTextModel(text_config)
|
|
||||||
self._uses_bitsandbytes_4bit = getattr(
|
self._uses_bitsandbytes_4bit = getattr(
|
||||||
arch_config, "ideogram_bnb_4bit_weight_only", False
|
arch_config, "ideogram_bnb_4bit_weight_only", False
|
||||||
)
|
)
|
||||||
|
self._uses_weight_only_fp8 = getattr(
|
||||||
|
arch_config, "ideogram_fp8_weight_only", False
|
||||||
|
)
|
||||||
|
quant_config = None
|
||||||
if self._uses_bitsandbytes_4bit:
|
if self._uses_bitsandbytes_4bit:
|
||||||
swap_linears_to_bitsandbytes_4bit(self.language_model)
|
source_quant_config = getattr(arch_config, "quantization_config")
|
||||||
elif getattr(arch_config, "ideogram_fp8_weight_only", False):
|
if isinstance(source_quant_config, dict):
|
||||||
swap_linears_to_weight_only_fp8(self.language_model)
|
quant_config_dict = source_quant_config
|
||||||
|
else:
|
||||||
|
quant_config_dict = source_quant_config.to_dict()
|
||||||
|
quant_config = BitsAndBytesConfig.from_config(quant_config_dict)
|
||||||
|
self.language_model = Qwen3VLTextModel(
|
||||||
|
text_config,
|
||||||
|
quant_config=quant_config,
|
||||||
|
use_weight_only_fp8=self._uses_weight_only_fp8,
|
||||||
|
use_tensor_parallel=True,
|
||||||
|
)
|
||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def forward(
|
def forward(
|
||||||
|
|||||||
@@ -9,7 +9,23 @@ from transformers.modeling_flash_attention_utils import FlashAttentionKwargs
|
|||||||
from transformers.utils import TransformersKwargs, is_torchdynamo_compiling
|
from transformers.utils import TransformersKwargs, is_torchdynamo_compiling
|
||||||
|
|
||||||
from sglang.multimodal_gen.configs.models.encoders.qwen3vl import Qwen3VLConfig
|
from sglang.multimodal_gen.configs.models.encoders.qwen3vl import Qwen3VLConfig
|
||||||
|
from sglang.multimodal_gen.runtime.distributed import (
|
||||||
|
get_tp_world_size,
|
||||||
|
model_parallel_is_initialized,
|
||||||
|
tensor_model_parallel_all_gather,
|
||||||
|
)
|
||||||
from sglang.multimodal_gen.runtime.layers.attention import LocalAttention
|
from sglang.multimodal_gen.runtime.layers.attention import LocalAttention
|
||||||
|
from sglang.multimodal_gen.runtime.layers.linear import (
|
||||||
|
ColumnParallelLinear,
|
||||||
|
ReplicatedLinear,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.layers.quantization.configs.base_config import (
|
||||||
|
QuantizationConfig,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.layers.quantization.weight_only_fp8 import (
|
||||||
|
WeightOnlyFP8ColumnParallelLinear,
|
||||||
|
WeightOnlyFP8Linear,
|
||||||
|
)
|
||||||
from sglang.multimodal_gen.runtime.loader.weight_utils import default_weight_loader
|
from sglang.multimodal_gen.runtime.loader.weight_utils import default_weight_loader
|
||||||
from sglang.multimodal_gen.runtime.models.encoders.base import TextEncoder
|
from sglang.multimodal_gen.runtime.models.encoders.base import TextEncoder
|
||||||
from sglang.multimodal_gen.runtime.platforms import AttentionBackendEnum
|
from sglang.multimodal_gen.runtime.platforms import AttentionBackendEnum
|
||||||
@@ -44,42 +60,152 @@ from transformers.models.qwen3_vl.modeling_qwen3_vl import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class Qwen3VLQuantizedLinear(ReplicatedLinear):
|
||||||
|
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
|
return super().forward(x)[0]
|
||||||
|
|
||||||
|
|
||||||
|
class Qwen3VLColumnParallelLinear(ColumnParallelLinear):
|
||||||
|
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
|
return super().forward(x)[0]
|
||||||
|
|
||||||
|
|
||||||
|
def _tp_world_size() -> int:
|
||||||
|
if not model_parallel_is_initialized():
|
||||||
|
return 1
|
||||||
|
return get_tp_world_size()
|
||||||
|
|
||||||
|
|
||||||
|
def _make_text_linear(
|
||||||
|
in_features: int,
|
||||||
|
out_features: int,
|
||||||
|
*,
|
||||||
|
bias: bool,
|
||||||
|
quant_config: QuantizationConfig | None,
|
||||||
|
use_weight_only_fp8: bool,
|
||||||
|
use_tensor_parallel: bool,
|
||||||
|
gather_output: bool,
|
||||||
|
prefix: str,
|
||||||
|
):
|
||||||
|
tp_size = _tp_world_size()
|
||||||
|
use_column_parallel = (
|
||||||
|
use_tensor_parallel and tp_size > 1 and out_features % tp_size == 0
|
||||||
|
)
|
||||||
|
if use_weight_only_fp8:
|
||||||
|
if use_column_parallel:
|
||||||
|
return WeightOnlyFP8ColumnParallelLinear(
|
||||||
|
in_features,
|
||||||
|
out_features,
|
||||||
|
bias=bias,
|
||||||
|
gather_output=gather_output,
|
||||||
|
)
|
||||||
|
return WeightOnlyFP8Linear(in_features, out_features, bias=bias)
|
||||||
|
if quant_config is not None:
|
||||||
|
if use_column_parallel:
|
||||||
|
return Qwen3VLColumnParallelLinear(
|
||||||
|
in_features,
|
||||||
|
out_features,
|
||||||
|
bias=bias,
|
||||||
|
gather_output=gather_output,
|
||||||
|
quant_config=quant_config,
|
||||||
|
prefix=prefix,
|
||||||
|
)
|
||||||
|
return Qwen3VLQuantizedLinear(
|
||||||
|
in_features,
|
||||||
|
out_features,
|
||||||
|
bias=bias,
|
||||||
|
quant_config=quant_config,
|
||||||
|
prefix=prefix,
|
||||||
|
)
|
||||||
|
if use_column_parallel:
|
||||||
|
return Qwen3VLColumnParallelLinear(
|
||||||
|
in_features,
|
||||||
|
out_features,
|
||||||
|
bias=bias,
|
||||||
|
gather_output=gather_output,
|
||||||
|
quant_config=None,
|
||||||
|
prefix=prefix,
|
||||||
|
)
|
||||||
|
return nn.Linear(in_features, out_features, bias=bias)
|
||||||
|
|
||||||
|
|
||||||
|
def _gather_tensor_parallel_activation(
|
||||||
|
x: torch.Tensor, linear: nn.Module
|
||||||
|
) -> torch.Tensor:
|
||||||
|
tp_group = getattr(linear, "tp_group", None)
|
||||||
|
if tp_group is None:
|
||||||
|
return x
|
||||||
|
return tensor_model_parallel_all_gather(x, tp_group=tp_group)
|
||||||
|
|
||||||
|
|
||||||
class Qwen3VLTextAttention(nn.Module):
|
class Qwen3VLTextAttention(nn.Module):
|
||||||
"""Multi-headed attention from 'Attention Is All You Need' paper"""
|
"""Multi-headed attention from 'Attention Is All You Need' paper"""
|
||||||
|
|
||||||
def __init__(self, config: Qwen3VLTextConfig, layer_idx: int):
|
def __init__(
|
||||||
|
self,
|
||||||
|
config: Qwen3VLTextConfig,
|
||||||
|
layer_idx: int,
|
||||||
|
quant_config: QuantizationConfig | None = None,
|
||||||
|
use_weight_only_fp8: bool = False,
|
||||||
|
use_tensor_parallel: bool = False,
|
||||||
|
prefix: str = "",
|
||||||
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.config = config
|
self.config = config
|
||||||
self.layer_idx = layer_idx
|
self.layer_idx = layer_idx
|
||||||
self.head_dim = config.hidden_size // config.num_attention_heads
|
self.head_dim = config.hidden_size // config.num_attention_heads
|
||||||
self.num_key_value_groups = (
|
self.total_num_heads = config.num_attention_heads
|
||||||
config.num_attention_heads // config.num_key_value_heads
|
self.total_num_key_value_heads = config.num_key_value_heads
|
||||||
)
|
self.tp_size = _tp_world_size() if use_tensor_parallel else 1
|
||||||
|
if self.tp_size > 1:
|
||||||
|
assert self.total_num_heads % self.tp_size == 0
|
||||||
|
assert self.total_num_key_value_heads % self.tp_size == 0
|
||||||
|
self.num_heads = self.total_num_heads // self.tp_size
|
||||||
|
self.num_key_value_heads = self.total_num_key_value_heads // self.tp_size
|
||||||
|
self.num_key_value_groups = self.num_heads // self.num_key_value_heads
|
||||||
self.scaling = self.head_dim**-0.5
|
self.scaling = self.head_dim**-0.5
|
||||||
self.attention_dropout = config.attention_dropout
|
self.attention_dropout = config.attention_dropout
|
||||||
self.is_causal = True
|
self.is_causal = True
|
||||||
self.num_heads = config.num_attention_heads
|
|
||||||
self.num_key_value_heads = config.num_key_value_heads
|
|
||||||
|
|
||||||
self.q_proj = nn.Linear(
|
self.q_proj = _make_text_linear(
|
||||||
config.hidden_size,
|
config.hidden_size,
|
||||||
config.num_attention_heads * self.head_dim,
|
config.num_attention_heads * self.head_dim,
|
||||||
bias=config.attention_bias,
|
bias=config.attention_bias,
|
||||||
|
quant_config=quant_config,
|
||||||
|
use_weight_only_fp8=use_weight_only_fp8,
|
||||||
|
use_tensor_parallel=use_tensor_parallel,
|
||||||
|
gather_output=False,
|
||||||
|
prefix=f"{prefix}.q_proj",
|
||||||
)
|
)
|
||||||
self.k_proj = nn.Linear(
|
self.k_proj = _make_text_linear(
|
||||||
config.hidden_size,
|
config.hidden_size,
|
||||||
config.num_key_value_heads * self.head_dim,
|
config.num_key_value_heads * self.head_dim,
|
||||||
bias=config.attention_bias,
|
bias=config.attention_bias,
|
||||||
|
quant_config=quant_config,
|
||||||
|
use_weight_only_fp8=use_weight_only_fp8,
|
||||||
|
use_tensor_parallel=use_tensor_parallel,
|
||||||
|
gather_output=False,
|
||||||
|
prefix=f"{prefix}.k_proj",
|
||||||
)
|
)
|
||||||
self.v_proj = nn.Linear(
|
self.v_proj = _make_text_linear(
|
||||||
config.hidden_size,
|
config.hidden_size,
|
||||||
config.num_key_value_heads * self.head_dim,
|
config.num_key_value_heads * self.head_dim,
|
||||||
bias=config.attention_bias,
|
bias=config.attention_bias,
|
||||||
|
quant_config=quant_config,
|
||||||
|
use_weight_only_fp8=use_weight_only_fp8,
|
||||||
|
use_tensor_parallel=use_tensor_parallel,
|
||||||
|
gather_output=False,
|
||||||
|
prefix=f"{prefix}.v_proj",
|
||||||
)
|
)
|
||||||
self.o_proj = nn.Linear(
|
self.o_proj = _make_text_linear(
|
||||||
config.num_attention_heads * self.head_dim,
|
config.num_attention_heads * self.head_dim,
|
||||||
config.hidden_size,
|
config.hidden_size,
|
||||||
bias=config.attention_bias,
|
bias=config.attention_bias,
|
||||||
|
quant_config=quant_config,
|
||||||
|
use_weight_only_fp8=use_weight_only_fp8,
|
||||||
|
use_tensor_parallel=use_tensor_parallel,
|
||||||
|
gather_output=True,
|
||||||
|
prefix=f"{prefix}.o_proj",
|
||||||
)
|
)
|
||||||
self.q_norm = Qwen3VLTextRMSNorm(
|
self.q_norm = Qwen3VLTextRMSNorm(
|
||||||
self.head_dim, eps=config.rms_norm_eps
|
self.head_dim, eps=config.rms_norm_eps
|
||||||
@@ -150,34 +276,94 @@ class Qwen3VLTextAttention(nn.Module):
|
|||||||
attn_output = self.attn(query_states, key_states, value_states)
|
attn_output = self.attn(query_states, key_states, value_states)
|
||||||
|
|
||||||
attn_output = attn_output.reshape(*input_shape, -1).contiguous()
|
attn_output = attn_output.reshape(*input_shape, -1).contiguous()
|
||||||
|
attn_output = _gather_tensor_parallel_activation(attn_output, self.q_proj)
|
||||||
attn_output = self.o_proj(attn_output)
|
attn_output = self.o_proj(attn_output)
|
||||||
return attn_output
|
return attn_output
|
||||||
|
|
||||||
|
|
||||||
class Qwen3VLTextMLP(nn.Module):
|
class Qwen3VLTextMLP(nn.Module):
|
||||||
def __init__(self, config):
|
def __init__(
|
||||||
|
self,
|
||||||
|
config,
|
||||||
|
quant_config: QuantizationConfig | None = None,
|
||||||
|
use_weight_only_fp8: bool = False,
|
||||||
|
use_tensor_parallel: bool = False,
|
||||||
|
prefix: str = "",
|
||||||
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.config = config
|
self.config = config
|
||||||
self.hidden_size = config.hidden_size
|
self.hidden_size = config.hidden_size
|
||||||
self.intermediate_size = config.intermediate_size
|
self.intermediate_size = config.intermediate_size
|
||||||
self.gate_proj = nn.Linear(self.hidden_size, self.intermediate_size, bias=False)
|
self.gate_proj = _make_text_linear(
|
||||||
self.up_proj = nn.Linear(self.hidden_size, self.intermediate_size, bias=False)
|
self.hidden_size,
|
||||||
self.down_proj = nn.Linear(self.intermediate_size, self.hidden_size, bias=False)
|
self.intermediate_size,
|
||||||
|
bias=False,
|
||||||
|
quant_config=quant_config,
|
||||||
|
use_weight_only_fp8=use_weight_only_fp8,
|
||||||
|
use_tensor_parallel=use_tensor_parallel,
|
||||||
|
gather_output=False,
|
||||||
|
prefix=f"{prefix}.gate_proj",
|
||||||
|
)
|
||||||
|
self.up_proj = _make_text_linear(
|
||||||
|
self.hidden_size,
|
||||||
|
self.intermediate_size,
|
||||||
|
bias=False,
|
||||||
|
quant_config=quant_config,
|
||||||
|
use_weight_only_fp8=use_weight_only_fp8,
|
||||||
|
use_tensor_parallel=use_tensor_parallel,
|
||||||
|
gather_output=False,
|
||||||
|
prefix=f"{prefix}.up_proj",
|
||||||
|
)
|
||||||
|
self.down_proj = _make_text_linear(
|
||||||
|
self.intermediate_size,
|
||||||
|
self.hidden_size,
|
||||||
|
bias=False,
|
||||||
|
quant_config=quant_config,
|
||||||
|
use_weight_only_fp8=use_weight_only_fp8,
|
||||||
|
use_tensor_parallel=use_tensor_parallel,
|
||||||
|
gather_output=True,
|
||||||
|
prefix=f"{prefix}.down_proj",
|
||||||
|
)
|
||||||
self.act_fn = ACT2FN[config.hidden_act]
|
self.act_fn = ACT2FN[config.hidden_act]
|
||||||
|
|
||||||
def forward(self, x):
|
def forward(self, x):
|
||||||
down_proj = self.down_proj(self.act_fn(self.gate_proj(x)) * self.up_proj(x))
|
hidden_states = self.act_fn(self.gate_proj(x)) * self.up_proj(x)
|
||||||
|
hidden_states = _gather_tensor_parallel_activation(
|
||||||
|
hidden_states, self.gate_proj
|
||||||
|
)
|
||||||
|
down_proj = self.down_proj(hidden_states)
|
||||||
return down_proj
|
return down_proj
|
||||||
|
|
||||||
|
|
||||||
class Qwen3VLTextDecoderLayer(nn.Module):
|
class Qwen3VLTextDecoderLayer(nn.Module):
|
||||||
def __init__(self, config: Qwen3VLTextConfig, layer_idx: int):
|
def __init__(
|
||||||
|
self,
|
||||||
|
config: Qwen3VLTextConfig,
|
||||||
|
layer_idx: int,
|
||||||
|
quant_config: QuantizationConfig | None = None,
|
||||||
|
use_weight_only_fp8: bool = False,
|
||||||
|
use_tensor_parallel: bool = False,
|
||||||
|
prefix: str = "",
|
||||||
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.hidden_size = config.hidden_size
|
self.hidden_size = config.hidden_size
|
||||||
|
|
||||||
self.self_attn = Qwen3VLTextAttention(config=config, layer_idx=layer_idx)
|
self.self_attn = Qwen3VLTextAttention(
|
||||||
|
config=config,
|
||||||
|
layer_idx=layer_idx,
|
||||||
|
quant_config=quant_config,
|
||||||
|
use_weight_only_fp8=use_weight_only_fp8,
|
||||||
|
use_tensor_parallel=use_tensor_parallel,
|
||||||
|
prefix=f"{prefix}.self_attn",
|
||||||
|
)
|
||||||
|
|
||||||
self.mlp = Qwen3VLTextMLP(config)
|
self.mlp = Qwen3VLTextMLP(
|
||||||
|
config,
|
||||||
|
quant_config=quant_config,
|
||||||
|
use_weight_only_fp8=use_weight_only_fp8,
|
||||||
|
use_tensor_parallel=use_tensor_parallel,
|
||||||
|
prefix=f"{prefix}.mlp",
|
||||||
|
)
|
||||||
self.input_layernorm = Qwen3VLTextRMSNorm(
|
self.input_layernorm = Qwen3VLTextRMSNorm(
|
||||||
config.hidden_size, eps=config.rms_norm_eps
|
config.hidden_size, eps=config.rms_norm_eps
|
||||||
)
|
)
|
||||||
@@ -223,7 +409,13 @@ class Qwen3VLTextModel(nn.Module):
|
|||||||
config: Qwen3VLTextConfig
|
config: Qwen3VLTextConfig
|
||||||
_no_split_modules = ["Qwen3VLTextDecoderLayer"]
|
_no_split_modules = ["Qwen3VLTextDecoderLayer"]
|
||||||
|
|
||||||
def __init__(self, config: Qwen3VLTextConfig):
|
def __init__(
|
||||||
|
self,
|
||||||
|
config: Qwen3VLTextConfig,
|
||||||
|
quant_config: QuantizationConfig | None = None,
|
||||||
|
use_weight_only_fp8: bool = False,
|
||||||
|
use_tensor_parallel: bool = False,
|
||||||
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.config = config
|
self.config = config
|
||||||
self.padding_idx = config.pad_token_id
|
self.padding_idx = config.pad_token_id
|
||||||
@@ -234,7 +426,14 @@ class Qwen3VLTextModel(nn.Module):
|
|||||||
)
|
)
|
||||||
self.layers = nn.ModuleList(
|
self.layers = nn.ModuleList(
|
||||||
[
|
[
|
||||||
Qwen3VLTextDecoderLayer(config, layer_idx)
|
Qwen3VLTextDecoderLayer(
|
||||||
|
config,
|
||||||
|
layer_idx,
|
||||||
|
quant_config=quant_config,
|
||||||
|
use_weight_only_fp8=use_weight_only_fp8,
|
||||||
|
use_tensor_parallel=use_tensor_parallel,
|
||||||
|
prefix=f"layers.{layer_idx}",
|
||||||
|
)
|
||||||
for layer_idx in range(config.num_hidden_layers)
|
for layer_idx in range(config.num_hidden_layers)
|
||||||
]
|
]
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -19,7 +19,7 @@
|
|||||||
"psnr_threshold": 28.0,
|
"psnr_threshold": 28.0,
|
||||||
"mean_abs_diff_threshold": 8.0
|
"mean_abs_diff_threshold": 8.0
|
||||||
},
|
},
|
||||||
"ideogram4_fp8_t2i": {
|
"ideogram4_nvfp4_t2i": {
|
||||||
"clip_threshold": 0.97,
|
"clip_threshold": 0.97,
|
||||||
"ssim_threshold": 0.78,
|
"ssim_threshold": 0.78,
|
||||||
"psnr_threshold": 18.0,
|
"psnr_threshold": 18.0,
|
||||||
|
|||||||
@@ -96,16 +96,6 @@ ONE_GPU_CASES: list[DiffusionTestCase] = [
|
|||||||
run_models_api_check=False,
|
run_models_api_check=False,
|
||||||
run_t2v_input_reference_check=False,
|
run_t2v_input_reference_check=False,
|
||||||
),
|
),
|
||||||
DiffusionTestCase(
|
|
||||||
"ideogram4_fp8_t2i",
|
|
||||||
DiffusionServerArgs(
|
|
||||||
model_path="ideogram-ai/ideogram-4-fp8",
|
|
||||||
),
|
|
||||||
IDEOGRAM4_CI_sampling_params,
|
|
||||||
run_perf_check=True,
|
|
||||||
run_consistency_check=True,
|
|
||||||
run_component_accuracy_check=False,
|
|
||||||
),
|
|
||||||
DiffusionTestCase(
|
DiffusionTestCase(
|
||||||
"flux_image_t2i",
|
"flux_image_t2i",
|
||||||
DiffusionServerArgs(model_path=DEFAULT_FLUX_1_DEV_MODEL_NAME_FOR_TEST),
|
DiffusionServerArgs(model_path=DEFAULT_FLUX_1_DEV_MODEL_NAME_FOR_TEST),
|
||||||
@@ -576,6 +566,21 @@ else:
|
|||||||
ONE_GPU_B200_CASES = ONE_GPU_MODELOPT_NVFP4_CASES
|
ONE_GPU_B200_CASES = ONE_GPU_MODELOPT_NVFP4_CASES
|
||||||
|
|
||||||
TWO_GPU_CASES = [
|
TWO_GPU_CASES = [
|
||||||
|
DiffusionTestCase(
|
||||||
|
"ideogram4_fp8_tp2_t2i",
|
||||||
|
DiffusionServerArgs(
|
||||||
|
model_path="ideogram-ai/ideogram-4-fp8",
|
||||||
|
tp_size=2,
|
||||||
|
extras=[
|
||||||
|
"--attention-backend",
|
||||||
|
"fa",
|
||||||
|
],
|
||||||
|
),
|
||||||
|
IDEOGRAM4_CI_sampling_params,
|
||||||
|
run_perf_check=False,
|
||||||
|
run_consistency_check=False,
|
||||||
|
run_component_accuracy_check=False,
|
||||||
|
),
|
||||||
DiffusionTestCase(
|
DiffusionTestCase(
|
||||||
"wan2_2_i2v_a14b_2gpu",
|
"wan2_2_i2v_a14b_2gpu",
|
||||||
DiffusionServerArgs(
|
DiffusionServerArgs(
|
||||||
|
|||||||
@@ -25,12 +25,16 @@ from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType, get_mod
|
|||||||
from sglang.multimodal_gen.runtime.distributed import get_local_torch_device
|
from sglang.multimodal_gen.runtime.distributed import get_local_torch_device
|
||||||
from sglang.multimodal_gen.runtime.layers.attention import USPAttention
|
from sglang.multimodal_gen.runtime.layers.attention import USPAttention
|
||||||
from sglang.multimodal_gen.runtime.layers.linear import UnquantizedLinearMethod
|
from sglang.multimodal_gen.runtime.layers.linear import UnquantizedLinearMethod
|
||||||
|
from sglang.multimodal_gen.runtime.layers.quantization.bitsandbytes import (
|
||||||
|
_maybe_shard_bitsandbytes_4bit_quant_state,
|
||||||
|
)
|
||||||
from sglang.multimodal_gen.runtime.layers.quantization.modelopt_quant import (
|
from sglang.multimodal_gen.runtime.layers.quantization.modelopt_quant import (
|
||||||
ModelOptFp4Config,
|
ModelOptFp4Config,
|
||||||
ModelOptFp4LinearMethod,
|
ModelOptFp4LinearMethod,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.layers.quantization.weight_only_fp8 import (
|
from sglang.multimodal_gen.runtime.layers.quantization.weight_only_fp8 import (
|
||||||
FP8_WEIGHT_DTYPE,
|
FP8_WEIGHT_DTYPE,
|
||||||
|
WeightOnlyFP8ColumnParallelLinear,
|
||||||
WeightOnlyFP8Linear,
|
WeightOnlyFP8Linear,
|
||||||
dequantize_rowwise_fp8_weight,
|
dequantize_rowwise_fp8_weight,
|
||||||
)
|
)
|
||||||
@@ -118,6 +122,29 @@ class FakeIdeogramPipeline:
|
|||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
|
class FakeBnbQuantState:
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
absmax,
|
||||||
|
shape=None,
|
||||||
|
code=None,
|
||||||
|
blocksize=None,
|
||||||
|
quant_type=None,
|
||||||
|
dtype=None,
|
||||||
|
offset=None,
|
||||||
|
state2=None,
|
||||||
|
):
|
||||||
|
self.absmax = absmax
|
||||||
|
self.shape = shape
|
||||||
|
self.code = code
|
||||||
|
self.blocksize = blocksize
|
||||||
|
self.quant_type = quant_type
|
||||||
|
self.dtype = dtype
|
||||||
|
self.offset = offset
|
||||||
|
self.state2 = state2
|
||||||
|
self.nested = state2 is not None
|
||||||
|
|
||||||
|
|
||||||
def _fake_server_args(cfg=None):
|
def _fake_server_args(cfg=None):
|
||||||
return SimpleNamespace(
|
return SimpleNamespace(
|
||||||
pipeline_config=cfg or Ideogram4PipelineConfig(),
|
pipeline_config=cfg or Ideogram4PipelineConfig(),
|
||||||
@@ -629,6 +656,48 @@ class TestIdeogram4(unittest.TestCase):
|
|||||||
)
|
)
|
||||||
self.assertEqual(state["layers.0.attention.qkv.weight"].dtype, FP8_WEIGHT_DTYPE)
|
self.assertEqual(state["layers.0.attention.qkv.weight"].dtype, FP8_WEIGHT_DTYPE)
|
||||||
|
|
||||||
|
def test_ideogram_dit_uses_tp_fp8_linears_when_tp_is_initialized(self):
|
||||||
|
import sglang.multimodal_gen.runtime.server_args as server_args_module
|
||||||
|
|
||||||
|
fake_tp_group = SimpleNamespace(world_size=2, rank_in_group=1)
|
||||||
|
prev_args = server_args_module._global_server_args
|
||||||
|
try:
|
||||||
|
set_global_server_args(
|
||||||
|
SimpleNamespace(attention_backend="torch_sdpa", comfyui_mode=False)
|
||||||
|
)
|
||||||
|
with (
|
||||||
|
patch(
|
||||||
|
"sglang.multimodal_gen.runtime.models.dits.ideogram.model_parallel_is_initialized",
|
||||||
|
return_value=True,
|
||||||
|
),
|
||||||
|
patch(
|
||||||
|
"sglang.multimodal_gen.runtime.models.dits.ideogram.get_tp_world_size",
|
||||||
|
return_value=2,
|
||||||
|
),
|
||||||
|
patch(
|
||||||
|
"sglang.multimodal_gen.runtime.layers.linear.get_tp_group",
|
||||||
|
return_value=fake_tp_group,
|
||||||
|
),
|
||||||
|
patch(
|
||||||
|
"sglang.multimodal_gen.runtime.layers.quantization.weight_only_fp8.get_tp_group",
|
||||||
|
return_value=fake_tp_group,
|
||||||
|
),
|
||||||
|
patch(
|
||||||
|
"sglang.multimodal_gen.runtime.layers.attention.layer.get_ring_parallel_world_size",
|
||||||
|
return_value=1,
|
||||||
|
),
|
||||||
|
):
|
||||||
|
with torch.device("meta"):
|
||||||
|
model = Ideogram4Transformer2DModel(Ideogram4DiTConfig(), {})
|
||||||
|
finally:
|
||||||
|
set_global_server_args(prev_args)
|
||||||
|
|
||||||
|
self.assertIsInstance(model.input_proj, WeightOnlyFP8ColumnParallelLinear)
|
||||||
|
self.assertEqual(tuple(model.input_proj.weight.shape), (2304, 128))
|
||||||
|
self.assertEqual(
|
||||||
|
tuple(model.layers[0].attention.qkv.weight.shape), (6912, 4608)
|
||||||
|
)
|
||||||
|
|
||||||
def test_ideogram_dit_nvfp4_quant_config_uses_native_fp4_linears(self):
|
def test_ideogram_dit_nvfp4_quant_config_uses_native_fp4_linears(self):
|
||||||
import sglang.multimodal_gen.runtime.server_args as server_args_module
|
import sglang.multimodal_gen.runtime.server_args as server_args_module
|
||||||
|
|
||||||
@@ -692,6 +761,104 @@ class TestIdeogram4(unittest.TestCase):
|
|||||||
(1,),
|
(1,),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def test_ideogram_dit_tp_nvfp4_uses_column_parallel_quant_linears(self):
|
||||||
|
import sglang.multimodal_gen.runtime.server_args as server_args_module
|
||||||
|
|
||||||
|
fake_tp_group = SimpleNamespace(world_size=2, rank_in_group=1)
|
||||||
|
quant_config = ModelOptFp4Config(
|
||||||
|
is_checkpoint_nvfp4_serialized=True,
|
||||||
|
group_size=16,
|
||||||
|
)
|
||||||
|
prev_args = server_args_module._global_server_args
|
||||||
|
try:
|
||||||
|
set_global_server_args(
|
||||||
|
SimpleNamespace(attention_backend="torch_sdpa", comfyui_mode=False)
|
||||||
|
)
|
||||||
|
with (
|
||||||
|
patch(
|
||||||
|
"sglang.multimodal_gen.runtime.models.dits.ideogram.model_parallel_is_initialized",
|
||||||
|
return_value=True,
|
||||||
|
),
|
||||||
|
patch(
|
||||||
|
"sglang.multimodal_gen.runtime.models.dits.ideogram.get_tp_world_size",
|
||||||
|
return_value=2,
|
||||||
|
),
|
||||||
|
patch(
|
||||||
|
"sglang.multimodal_gen.runtime.layers.linear.get_tp_group",
|
||||||
|
return_value=fake_tp_group,
|
||||||
|
),
|
||||||
|
patch(
|
||||||
|
"sglang.multimodal_gen.runtime.layers.attention.layer.get_ring_parallel_world_size",
|
||||||
|
return_value=1,
|
||||||
|
),
|
||||||
|
):
|
||||||
|
with torch.device("meta"):
|
||||||
|
model = Ideogram4Transformer2DModel(
|
||||||
|
Ideogram4DiTConfig(),
|
||||||
|
{},
|
||||||
|
quant_config=quant_config,
|
||||||
|
)
|
||||||
|
finally:
|
||||||
|
set_global_server_args(prev_args)
|
||||||
|
|
||||||
|
self.assertTrue(model.layers[0].attention.qkv.gather_output)
|
||||||
|
self.assertEqual(
|
||||||
|
tuple(model.layers[0].attention.qkv.weight.shape), (6912, 2304)
|
||||||
|
)
|
||||||
|
self.assertIsInstance(
|
||||||
|
model.layers[0].attention.qkv.quant_method,
|
||||||
|
ModelOptFp4LinearMethod,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_bitsandbytes_tp_quant_state_uses_local_output_shard(self):
|
||||||
|
param = torch.nn.Parameter(
|
||||||
|
torch.empty(8, 1, dtype=torch.uint8), requires_grad=False
|
||||||
|
)
|
||||||
|
param.bnb_full_shape = (4, 8)
|
||||||
|
param.bnb_local_shape = (2, 8)
|
||||||
|
param.bnb_output_shard_start = 2
|
||||||
|
param.bnb_input_shard_start = 0
|
||||||
|
quant_state = FakeBnbQuantState(
|
||||||
|
absmax=torch.arange(8, dtype=torch.float32),
|
||||||
|
shape=torch.Size((4, 8)),
|
||||||
|
code=torch.ones(16, dtype=torch.float32),
|
||||||
|
blocksize=4,
|
||||||
|
quant_type="nf4",
|
||||||
|
dtype=torch.bfloat16,
|
||||||
|
)
|
||||||
|
|
||||||
|
sharded = _maybe_shard_bitsandbytes_4bit_quant_state(param, quant_state)
|
||||||
|
|
||||||
|
self.assertEqual(sharded.shape, torch.Size((2, 8)))
|
||||||
|
torch.testing.assert_close(sharded.absmax, torch.tensor([4.0, 5.0, 6.0, 7.0]))
|
||||||
|
|
||||||
|
def test_assign_load_preserves_bitsandbytes_tp_attrs(self):
|
||||||
|
class TinyModule(torch.nn.Module):
|
||||||
|
def __init__(self):
|
||||||
|
super().__init__()
|
||||||
|
self.weight = torch.nn.Parameter(
|
||||||
|
torch.empty(8, 1, dtype=torch.uint8), requires_grad=False
|
||||||
|
)
|
||||||
|
self.weight.bnb_full_shape = (4, 8)
|
||||||
|
self.weight.bnb_local_shape = (2, 8)
|
||||||
|
self.weight.bnb_output_shard_start = 2
|
||||||
|
self.weight.bnb_input_shard_start = 0
|
||||||
|
|
||||||
|
model = TinyModule()
|
||||||
|
load_model_from_full_model_state_dict(
|
||||||
|
model,
|
||||||
|
iter([("weight", torch.ones(8, 1, dtype=torch.uint8))]),
|
||||||
|
torch.device("cpu"),
|
||||||
|
param_dtype=None,
|
||||||
|
strict=True,
|
||||||
|
param_names_mapping=lambda name: (name, None, None),
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(model.weight.bnb_full_shape, (4, 8))
|
||||||
|
self.assertEqual(model.weight.bnb_local_shape, (2, 8))
|
||||||
|
self.assertEqual(model.weight.bnb_output_shard_start, 2)
|
||||||
|
self.assertEqual(model.weight.bnb_input_shard_start, 0)
|
||||||
|
|
||||||
def test_missing_weight_only_fp8_scale_is_fatal(self):
|
def test_missing_weight_only_fp8_scale_is_fatal(self):
|
||||||
with torch.device("meta"):
|
with torch.device("meta"):
|
||||||
model = WeightOnlyFP8Linear(3, 2, bias=False)
|
model = WeightOnlyFP8Linear(3, 2, bias=False)
|
||||||
@@ -805,6 +972,57 @@ class TestIdeogram4(unittest.TestCase):
|
|||||||
any(isinstance(module, torch.nn.Linear) for module in encoder.modules())
|
any(isinstance(module, torch.nn.Linear) for module in encoder.modules())
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def test_ideogram_text_encoder_tp_fp8_uses_column_parallel_linears(self):
|
||||||
|
config = Ideogram4TextEncoderConfig()
|
||||||
|
config.post_diffusers_config_update()
|
||||||
|
config.arch_config.text_config = Qwen3VLTextConfig(
|
||||||
|
vocab_size=32,
|
||||||
|
hidden_size=16,
|
||||||
|
intermediate_size=32,
|
||||||
|
num_hidden_layers=1,
|
||||||
|
num_attention_heads=4,
|
||||||
|
num_key_value_heads=4,
|
||||||
|
head_dim=4,
|
||||||
|
max_position_embeddings=64,
|
||||||
|
pad_token_id=0,
|
||||||
|
)
|
||||||
|
import sglang.multimodal_gen.runtime.server_args as server_args_module
|
||||||
|
|
||||||
|
fake_tp_group = SimpleNamespace(world_size=2, rank_in_group=1)
|
||||||
|
prev_args = server_args_module._global_server_args
|
||||||
|
try:
|
||||||
|
set_global_server_args(
|
||||||
|
SimpleNamespace(attention_backend="torch_sdpa", comfyui_mode=False)
|
||||||
|
)
|
||||||
|
with (
|
||||||
|
patch(
|
||||||
|
"sglang.multimodal_gen.runtime.models.encoders.qwen3vl.model_parallel_is_initialized",
|
||||||
|
return_value=True,
|
||||||
|
),
|
||||||
|
patch(
|
||||||
|
"sglang.multimodal_gen.runtime.models.encoders.qwen3vl.get_tp_world_size",
|
||||||
|
return_value=2,
|
||||||
|
),
|
||||||
|
patch(
|
||||||
|
"sglang.multimodal_gen.runtime.layers.quantization.weight_only_fp8.get_tp_group",
|
||||||
|
return_value=fake_tp_group,
|
||||||
|
),
|
||||||
|
):
|
||||||
|
with torch.device("meta"):
|
||||||
|
encoder = IdeogramQwen3VLTextEncoder(config)
|
||||||
|
finally:
|
||||||
|
set_global_server_args(prev_args)
|
||||||
|
|
||||||
|
layer = encoder.language_model.layers[0]
|
||||||
|
self.assertEqual(layer.self_attn.num_heads, 2)
|
||||||
|
self.assertEqual(layer.self_attn.num_key_value_heads, 2)
|
||||||
|
self.assertIsInstance(layer.self_attn.q_proj, WeightOnlyFP8ColumnParallelLinear)
|
||||||
|
self.assertFalse(layer.self_attn.q_proj.gather_output)
|
||||||
|
self.assertTrue(layer.self_attn.o_proj.gather_output)
|
||||||
|
self.assertIsInstance(layer.mlp.gate_proj, WeightOnlyFP8ColumnParallelLinear)
|
||||||
|
self.assertFalse(layer.mlp.gate_proj.gather_output)
|
||||||
|
self.assertTrue(layer.mlp.down_proj.gather_output)
|
||||||
|
|
||||||
def test_denoise_and_decode_shape_smoke(self):
|
def test_denoise_and_decode_shape_smoke(self):
|
||||||
import sglang.multimodal_gen.runtime.server_args as server_args_module
|
import sglang.multimodal_gen.runtime.server_args as server_args_module
|
||||||
|
|
||||||
|
|||||||
@@ -562,9 +562,14 @@ class EagleDraftWorker(BaseDraftWorker):
|
|||||||
if (c := self.draft_runner.canary_manager) is not None
|
if (c := self.draft_runner.canary_manager) is not None
|
||||||
else contextlib.nullcontext()
|
else contextlib.nullcontext()
|
||||||
)
|
)
|
||||||
with forward_context(
|
with (
|
||||||
ForwardContext(attn_backend=self.draft_attn_backend.attn_backends[i])
|
forward_context(
|
||||||
), canary_index_ctx:
|
ForwardContext(
|
||||||
|
attn_backend=self.draft_attn_backend.attn_backends[i]
|
||||||
|
)
|
||||||
|
),
|
||||||
|
canary_index_ctx,
|
||||||
|
):
|
||||||
logits_output = self.draft_runner.forward(forward_batch).logits_output
|
logits_output = self.draft_runner.forward(forward_batch).logits_output
|
||||||
maybe_detect_nan(logits_output.next_token_logits, f"draft_forward step {i}")
|
maybe_detect_nan(logits_output.next_token_logits, f"draft_forward step {i}")
|
||||||
maybe_detect_inf(logits_output.next_token_logits, f"draft_forward step {i}")
|
maybe_detect_inf(logits_output.next_token_logits, f"draft_forward step {i}")
|
||||||
|
|||||||
Reference in New Issue
Block a user