Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions .github/workflows/integration_test_8gpu_rl.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@ on:
paths:
- 'torchtitan/experiments/rl/**'
- 'torchtitan/models/qwen3_5/**'
- 'torchtitan/models/qwen3_8/**'
- '.github/workflows/integration_test_8gpu_rl.yaml'
pull_request:
types: [labeled, synchronize]
Expand Down
30 changes: 17 additions & 13 deletions tests/unit_tests/cpu/components/data/test_qwen_multimodal_data.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,7 @@
from torchtitan.hf_datasets.multimodal.utils.image import resize_to_navit_patch_grid
from torchtitan.models.kimi_k2_7 import config_registry as kimi_configs
from torchtitan.models.qwen3_5 import config_registry as qwen35_configs
from torchtitan.models.qwen3_8 import config_registry as qwen38_configs


class _Tokenizer:
Expand Down Expand Up @@ -151,21 +152,24 @@ def test_kimi_multimodal_recipe_copies_unpacked_dataset(recipe_name):


@pytest.mark.parametrize(
"recipe_name",
("config_registry", "recipe_name"),
[
"qwen35_debugmodel",
"qwen35_debugmodel_moe",
"qwen35_0_8b",
"qwen35_2b",
"qwen35_4b",
"qwen35_9b",
"qwen35_27b",
"qwen35_35b_a3b",
"qwen35_122b_a10b",
"qwen35_397b_a17b",
(qwen35_configs, "qwen35_debugmodel"),
(qwen35_configs, "qwen35_debugmodel_moe"),
(qwen35_configs, "qwen35_0_8b"),
(qwen35_configs, "qwen35_2b"),
(qwen35_configs, "qwen35_4b"),
(qwen35_configs, "qwen35_9b"),
(qwen35_configs, "qwen35_27b"),
(qwen35_configs, "qwen35_35b_a3b"),
(qwen35_configs, "qwen35_122b_a10b"),
(qwen35_configs, "qwen35_397b_a17b"),
(qwen38_configs, "qwen38_debugmodel"),
(qwen38_configs, "qwen38_debugmodel_moe"),
(qwen38_configs, "qwen38_27b"),
],
)
def test_qwen35_recipe_geometry_matches_dataset_processor(recipe_name):
def test_qwen_recipe_geometry_matches_dataset_processor(config_registry, recipe_name):
registry_state = {
name: (
id(dataset),
Expand All @@ -178,7 +182,7 @@ def test_qwen35_recipe_geometry_matches_dataset_processor(recipe_name):
if isinstance(dataset.processor, MultiModalProcessor.Config)
}

config = getattr(qwen35_configs, recipe_name)()
config = getattr(config_registry, recipe_name)()
dataset = config.dataloader.dataset
collator = config.dataloader.collator

Expand Down
2 changes: 2 additions & 0 deletions tests/unit_tests/cpu/test_no_new_cli_options.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,7 @@
import warnings

import tyro

from torchtitan.trainer import Trainer

_FROZEN_CLI_OPTIONS = frozenset(
Expand Down Expand Up @@ -363,6 +364,7 @@ def _declared_cli_options(
("deepseek_v3", "deepseek_v3_debugmodel"),
("qwen3", "qwen3_debugmodel"),
("qwen3_5", "qwen35_debugmodel_moe"),
("qwen3_8", "qwen38_debugmodel_moe"),
("gpt_oss", "gpt_oss_debugmodel"),
("flux", "flux_debugmodel"),
("kimi_k2_7", "kimi_k2_5_debugmodel"),
Expand Down
81 changes: 81 additions & 0 deletions tests/unit_tests/cpu/test_qwen3_5.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,81 @@
# Copyright (c) Meta Platforms, Inc. and affiliates.
# All rights reserved.
#
# This source code is licensed under the BSD-style license found in the
# LICENSE file in the root directory of this source tree.

from typing import cast

import pytest

pytest.importorskip("fla")

from torchtitan.models.qwen3_5 import model_registry, Qwen35Model, qwen3_5_configs
from torchtitan.models.qwen3_5.config_registry import qwen35_0_8b, qwen35_27b
from torchtitan.models.qwen3_8 import model_registry as qwen3_8_model_registry


def test_qwen35_registry_keeps_released_flavors() -> None:
assert set(qwen3_5_configs) == {
"debugmodel",
"debugmodel_moe",
"0.8B",
"2B",
"4B",
"9B",
"27B",
"35B-A3B",
"122B-A10B",
"397B-A17B",
}


@pytest.mark.parametrize("flavor", sorted(qwen3_5_configs))
def test_qwen35_registry_builds_every_flavor(flavor: str) -> None:
model_spec = model_registry(
flavor,
moe_comm_backend=(
"standard" if flavor == "debugmodel_moe" or "-A" in flavor else None
),
)

assert model_spec.name == "qwen3_5"
assert model_spec.flavor == flavor


def test_qwen35_is_the_shared_model_implementation() -> None:
model_spec = model_registry("0.8B")
config = cast(Qwen35Model.Config, model_spec.model)
qwen38_config = qwen3_8_model_registry("27B").model

assert model_spec.name == "qwen3_5"
assert model_spec.flavor == "0.8B"
assert config.dim == 1024
assert len(config.layers) == 24
assert isinstance(qwen38_config, Qwen35Model.Config)


def test_qwen35_keeps_small_dense_and_moe_models() -> None:
dense_config = cast(Qwen35Model.Config, model_registry("0.8B").model)
moe_config = cast(
Qwen35Model.Config,
model_registry("35B-A3B", moe_comm_backend="standard").model,
)

assert dense_config.dim == 1024
assert moe_config.dim == 2048
assert moe_config.layers[0].moe is not None
assert moe_config.layers[0].moe.router.num_experts == 256
assert moe_config.layers[0].moe.router.top_k == 8


def test_qwen35_recipes_keep_versioned_hugging_face_paths() -> None:
small_config = qwen35_0_8b()
large_config = qwen35_27b()

assert small_config.hf_assets_path.endswith("Qwen3.5-0.8B")
assert small_config.model_spec is not None
assert small_config.model_spec.name == "qwen3_5"
assert large_config.hf_assets_path.endswith("Qwen3.5-27B")
assert large_config.model_spec is not None
assert large_config.model_spec.name == "qwen3_5"
188 changes: 188 additions & 0 deletions tests/unit_tests/cpu/test_qwen3_8.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,188 @@
# Copyright (c) Meta Platforms, Inc. and affiliates.
# All rights reserved.
#
# This source code is licensed under the BSD-style license found in the
# LICENSE file in the root directory of this source tree.

from dataclasses import replace
from typing import cast

import pytest
import torch

pytest.importorskip("fla")

from torchtitan.models.qwen3_5 import Qwen35Model, Qwen35StateDictAdapter
from torchtitan.models.qwen3_5.sharding import set_qwen35_sharding_config
from torchtitan.models.qwen3_8 import model_registry, qwen3_8_configs
from torchtitan.models.qwen3_8.config_registry import qwen38_27b, qwen38_2_4t_a95b


def test_qwen38_registry_exposes_only_qwen38_flavors() -> None:
assert set(qwen3_8_configs) == {
"debugmodel",
"debugmodel_moe",
"27B",
"2.4T-A95B",
}
for legacy_flavor in (
"0.8B",
"2B",
"4B",
"9B",
"35B-A3B",
"122B-A10B",
"397B-A17B",
):
with pytest.raises(KeyError):
model_registry(legacy_flavor)


def test_qwen38_27b_reuses_qwen35_multimodal_architecture() -> None:
model_spec = model_registry("27B")
config = cast(Qwen35Model.Config, model_spec.model)

assert model_spec.name == "qwen3_8"
assert config.dim == 5120
assert len(config.layers) == 64
assert config.vision_encoder is not None
assert config.vision_encoder.merger.fc2.out_features == 5120


def test_qwen38_recipes_use_released_hugging_face_paths() -> None:
dense_config = qwen38_27b()
moe_config = qwen38_2_4t_a95b()

assert dense_config.hf_assets_path.endswith("Qwen3.8-27B")
assert dense_config.model_spec is not None
assert dense_config.model_spec.name == "qwen3_8"
assert moe_config.hf_assets_path.endswith("Qwen3.8-2.4T-A95B")
assert moe_config.model_spec is not None
assert moe_config.model_spec.name == "qwen3_8"


def test_qwen38_2_4t_a95b_matches_hugging_face_config() -> None:
config = qwen3_8_configs["2.4T-A95B"](
attn_backend="flex",
moe_comm_backend="standard",
)

assert config.dim == 8192
assert len(config.layers) == 92
assert config.vision_encoder is None

linear_layer = config.layers[0]
assert linear_layer.delta_net is not None
assert linear_layer.delta_net.in_proj_q.out_features == 16 * 128
assert linear_layer.delta_net.in_proj_v.out_features == 128 * 128

full_attention_layer = config.layers[3]
assert full_attention_layer.attention is not None
assert full_attention_layer.attention.n_heads == 64
assert full_attention_layer.attention.n_kv_heads == 4

assert linear_layer.moe is not None
assert linear_layer.moe.router.num_experts == 512
assert linear_layer.moe.router.top_k == 10


def test_text_only_qwen38_sharding_does_not_require_vision() -> None:
config = qwen3_8_configs["2.4T-A95B"](
attn_backend="flex",
moe_comm_backend="standard",
)

set_qwen35_sharding_config(config, enable_sp=True, enable_ep=True)

assert config.tok_embeddings.sharding_config is not None
assert config.layers[0].sharding_config is not None


def test_shared_model_builds_without_vision_encoder() -> None:
config = qwen3_8_configs["debugmodel"](attn_backend="flex")
config = replace(
config,
vocab_size=128,
tok_embeddings=replace(config.tok_embeddings, num_embeddings=128),
lm_head=replace(config.lm_head, out_features=128),
vision_encoder=None,
)

model = config.build()

assert model.vision_encoder is None


def test_text_only_checkpoint_adapter_uses_model_prefix() -> None:
config = qwen3_8_configs["2.4T-A95B"](
attn_backend="flex",
moe_comm_backend="standard",
)
adapter = Qwen35StateDictAdapter(config, hf_assets_path=None)
embedding = torch.randn(2, 3)
lm_head = torch.randn(2, 3)

converted = adapter.from_hf(
{
"model.embed_tokens.weight": embedding,
"lm_head.weight": lm_head,
}
)
assert set(converted) == {"tok_embeddings.weight", "lm_head.weight"}
torch.testing.assert_close(converted["tok_embeddings.weight"], embedding)
torch.testing.assert_close(converted["lm_head.weight"], lm_head)

restored = adapter.to_hf(converted)
assert set(restored) == {"model.embed_tokens.weight", "lm_head.weight"}
torch.testing.assert_close(restored["model.embed_tokens.weight"], embedding)
torch.testing.assert_close(restored["lm_head.weight"], lm_head)


def test_multimodal_checkpoint_adapter_keeps_language_model_prefix() -> None:
config = qwen3_8_configs["27B"](attn_backend="flex")
adapter = Qwen35StateDictAdapter(config, hf_assets_path=None)
embedding = torch.randn(2, 3)
lm_head = torch.randn(2, 3)

converted = adapter.from_hf(
{
"model.language_model.embed_tokens.weight": embedding,
"lm_head.weight": lm_head,
}
)
restored = adapter.to_hf(converted)

assert "model.language_model.embed_tokens.weight" in restored
assert "model.embed_tokens.weight" not in restored


def test_text_only_checkpoint_adapter_converts_fused_deltanet_qkv() -> None:
config = qwen3_8_configs["2.4T-A95B"](
attn_backend="flex",
moe_comm_backend="standard",
)
adapter = Qwen35StateDictAdapter(config, hf_assets_path=None)
delta_net = config.layers[0].delta_net
assert delta_net is not None
key_dim = delta_net.in_proj_q.out_features
value_dim = delta_net.in_proj_v.out_features
fused_qkv = torch.randn(key_dim * 2 + value_dim, 1)

converted = adapter.from_hf(
{
"model.layers.0.linear_attn.in_proj_qkv.weight": fused_qkv,
"lm_head.weight": torch.randn(2, 3),
}
)
assert set(converted) == {
"layers.0.attn.in_proj_q.weight",
"layers.0.attn.in_proj_k.weight",
"layers.0.attn.in_proj_v.weight",
"lm_head.weight",
}

restored = adapter.to_hf(converted)
torch.testing.assert_close(
restored["model.layers.0.linear_attn.in_proj_qkv.weight"],
fused_qkv,
)
1 change: 1 addition & 0 deletions torchtitan/models/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,5 +15,6 @@
"muse_glimmer",
"qwen3",
"qwen3_5",
"qwen3_8",
]
)
14 changes: 11 additions & 3 deletions torchtitan/models/qwen3_5/model.py
Original file line number Diff line number Diff line change
Expand Up @@ -284,7 +284,7 @@ class Qwen35Model(Decoder):

@dataclass(kw_only=True, slots=True)
class Config(Decoder.Config):
vision_encoder: Qwen35VisionEncoder.Config
vision_encoder: Qwen35VisionEncoder.Config | None = None

def update_from_config(
self,
Expand Down Expand Up @@ -358,8 +358,14 @@ def get_nparams_and_flops(
def __init__(self, config: Config):
super().__init__(config)

self.vision_encoder = config.vision_encoder.build()
self.spatial_merge_size = config.vision_encoder.spatial_merge_size
self.vision_encoder = (

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Qwen3.8-2.4T-A95B model is text only so we want to allow None vision_encoder

config.vision_encoder.build() if config.vision_encoder is not None else None
)
self.spatial_merge_size = (
config.vision_encoder.spatial_merge_size
if config.vision_encoder is not None
else None
)

def preprocess_inputs(
self,
Expand Down Expand Up @@ -492,6 +498,8 @@ def _get_vision_embeds(
vision_embeds: Packed vision embeddings ``(total_tokens, dim)``.
num_tokens_per_item: (num_items,) actual token count per item
"""
if self.vision_encoder is None:
raise ValueError("Vision inputs were provided without a vision encoder.")
pixel_values = pixel_values.to(self.vision_encoder.patch_embed.weight.dtype)
vision_embeds = self.vision_encoder(pixel_values, grid_thw=grid_thw)

Expand Down
Loading
Loading