Skip to content
Merged
Show file tree
Hide file tree
Changes from 4 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,
)
Loading
Loading