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
94 changes: 93 additions & 1 deletion tests/utils/test_flops_counter.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,7 +18,18 @@

from verl.utils.flops_counter import FlopsCounter

VALID_CONFIG_TYPE = {"llama", "qwen2", "qwen3", "qwen3_moe", "deepseek_v3", "mistral", "gemma3_text", "apertus"}
VALID_CONFIG_TYPE = {
"llama",
"qwen2",
"qwen3",
"qwen3_moe",
"qwen3_5",
"qwen3_5_moe",
"deepseek_v3",
"mistral",
"gemma3_text",
"apertus",
}


class Config:
Expand Down Expand Up @@ -302,6 +313,85 @@ def __init__(self, config_dict):
# S*(2*V*H + L*(4*H**2 + k_mlp*H*I + k_qkn*H)) * (SUM[seqlen]) + 6*SUM[seqlen**2]*L*H
"expected_flops_tuple": (194825353691136 / 1e12, 692711652851712 / 1e12),
},
"qwen3_5": {
"config": { # Qwen/Qwen3.5-27B
"model_type": "qwen3_5",
"text_config": {
"vocab_size": 248320,
"hidden_size": 4096,
"intermediate_size": 12288,
"num_hidden_layers": 32,
"num_attention_heads": 16,
"num_key_value_heads": 4,
"head_dim": 256,
"linear_conv_kernel_dim": 4,
"linear_key_head_dim": 128,
"linear_value_head_dim": 128,
"linear_num_key_heads": 16,
"linear_num_value_heads": 32,
"layer_types": ["linear_attention" if bool((i + 1) % 4) else "full_attention" for i in range(32)],
},
"vision_config": {
"num_heads": 16,
"depth": 27,
"hidden_size": 1152,
"intermediate_size": 4304,
"out_hidden_size": 4096,
"spatial_merge_size": 2,
"temporal_patch_size": 2,
"in_channels": 3,
"patch_size": 16,
},
},
"batch_seqlens_tuple": ([512, 1024, 2048], [4096, 4096, 4096]),
"images_seqlens_tuple": ([512, 1024, 2048], [4096, 4096, 4096]),
# Text FLOPs include dense MLP, hybrid full attention/GatedDeltaNet projections, embeddings/lm_head,
# full-attention quadratic terms, and GatedDeltaNet recurrence. ViT FLOPs reuse Qwen3-VL accounting.
"expected_flops_tuple": (
206090394402816 / 1e12,
724521757704192 / 1e12,
),
},
"qwen3_5_moe": {
"config": { # Qwen/Qwen3.5-35B-A3B
"model_type": "qwen3_5_moe",
"text_config": {
"vocab_size": 248320,
"hidden_size": 2048,
"num_hidden_layers": 40,
"num_attention_heads": 16,
"num_key_value_heads": 2,
"head_dim": 256,
"linear_conv_kernel_dim": 4,
"linear_key_head_dim": 128,
"linear_value_head_dim": 128,
"linear_num_key_heads": 16,
"linear_num_value_heads": 32,
"moe_intermediate_size": 512,
"shared_expert_intermediate_size": 512,
"num_experts_per_tok": 8,
"num_experts": 256,
"layer_types": ["linear_attention" if bool((i + 1) % 4) else "full_attention" for i in range(40)],
},
"vision_config": {
"num_heads": 16,
"depth": 27,
"hidden_size": 1152,
"intermediate_size": 4304,
"out_hidden_size": 2048,
"spatial_merge_size": 2,
"temporal_patch_size": 2,
"in_channels": 3,
"patch_size": 16,
},
},
"batch_seqlens_tuple": ([512, 1024, 2048], [4096, 4096, 4096]),
"images_seqlens_tuple": ([512, 1024, 2048], [4096, 4096, 4096]),
"expected_flops_tuple": (
88082762170368 / 1e12,
321470349705216 / 1e12,
),
},
"qwen3_vl": {
"config": { # Qwen/Qwen3-VL-8B
"model_type": "qwen3_vl",
Expand Down Expand Up @@ -447,6 +537,8 @@ def __init__(self, config_dict):
"gemma3_text",
"apertus",
"gpt_oss",
"qwen3_5",
"qwen3_5_moe",
"qwen3_vl",
"qwen3_vl_moe",
],
Expand Down
94 changes: 94 additions & 0 deletions verl/utils/flops_counter.py
Original file line number Diff line number Diff line change
Expand Up @@ -265,6 +265,98 @@ def _estimate_qwen3_vit_flop(images_seqlens, config):
return vit_flops


def _count_qwen3_5_layer_types(config):
layer_types = getattr(config, "layer_types", None)
if layer_types:
num_full_attn_layers = sum(layer_type == "full_attention" for layer_type in layer_types)
num_linear_attn_layers = sum(layer_type == "linear_attention" for layer_type in layer_types)
return num_full_attn_layers, num_linear_attn_layers

full_attention_interval = getattr(config, "full_attention_interval", 4)
num_full_attn_layers = sum(
not bool((layer_idx + 1) % full_attention_interval) for layer_idx in range(config.num_hidden_layers)
)
return num_full_attn_layers, config.num_hidden_layers - num_full_attn_layers


def _compute_qwen3_5_hybrid_attn_params(config):
hidden_size = config.hidden_size
num_attention_heads = config.num_attention_heads
num_key_value_heads = config.num_key_value_heads
head_dim = getattr(config, "head_dim", hidden_size // num_attention_heads)

q_size = num_attention_heads * head_dim
k_size = num_key_value_heads * head_dim
v_size = num_key_value_heads * head_dim

num_full_attn_layers, num_linear_attn_layers = _count_qwen3_5_layer_types(config)

# Qwen3.5 full attention q_proj also emits a sigmoid gate, so q_proj is 2x the query size.
full_attn_linear_N = hidden_size * (2 * q_size + k_size + v_size + q_size)

linear_k_size = config.linear_num_key_heads * config.linear_key_head_dim
linear_v_size = config.linear_num_value_heads * config.linear_value_head_dim
linear_attn_linear_N = hidden_size * (2 * linear_k_size + 3 * linear_v_size + 2 * config.linear_num_value_heads)
conv_N = config.linear_conv_kernel_dim * (2 * linear_k_size + linear_v_size)

attn_linear_N = full_attn_linear_N * num_full_attn_layers
attn_linear_N += (linear_attn_linear_N + conv_N) * num_linear_attn_layers

return attn_linear_N, num_full_attn_layers, num_linear_attn_layers, head_dim, num_attention_heads


def _compute_qwen3_5_gdn_recurrence_flops(config, tokens_sum, num_linear_attn_layers):
return (
15
* config.linear_key_head_dim
* config.linear_value_head_dim
* config.linear_num_value_heads
* tokens_sum
* num_linear_attn_layers
)


def _estimate_qwen3_5_flops(config, tokens_sum, batch_seqlens, delta_time, **kargs):
# qwen3_5 and qwen3_5_moe use text_config and vision_config for the LLM and ViT parts.
text_config = config.text_config if hasattr(config, "text_config") else config
hidden_size = text_config.hidden_size
vocab_size = text_config.vocab_size
num_hidden_layers = text_config.num_hidden_layers

attn_linear_N, num_full_attn_layers, num_linear_attn_layers, head_dim, num_attention_heads = (
_compute_qwen3_5_hybrid_attn_params(text_config)
)

if hasattr(text_config, "num_experts"):
moe_gate_N = hidden_size * text_config.num_experts
moe_expertmlp_N = hidden_size * text_config.moe_intermediate_size * text_config.num_experts_per_tok * 3
moe_sharedexpertmlp_N = hidden_size * text_config.shared_expert_intermediate_size * 3
moe_sharedexpert_gate_N = hidden_size
mlp_N = (moe_gate_N + moe_expertmlp_N + moe_sharedexpertmlp_N + moe_sharedexpert_gate_N) * num_hidden_layers
else:
mlp_N = hidden_size * text_config.intermediate_size * 3 * num_hidden_layers

emd_and_lm_head_N = vocab_size * hidden_size * 2
dense_N_flops = 6 * (mlp_N + attn_linear_N + emd_and_lm_head_N) * tokens_sum

seqlen_square_sum = 0
for seqlen in batch_seqlens:
seqlen_square_sum += seqlen * seqlen
attn_qkv_flops = 6 * seqlen_square_sum * head_dim * num_attention_heads * num_full_attn_layers

gdn_recurrence_flops = _compute_qwen3_5_gdn_recurrence_flops(text_config, tokens_sum, num_linear_attn_layers)

images_seqlens = kargs.get("images_seqlens", None)
if images_seqlens is not None:
vit_flops = _estimate_qwen3_vit_flop(images_seqlens, config.vision_config)
else:
vit_flops = 0

flops_all_token = dense_N_flops + attn_qkv_flops + gdn_recurrence_flops + vit_flops
flops_achieved = flops_all_token * (1.0 / delta_time) / 1e12
return flops_achieved


def _estimate_deepseek_v3_flops(config, tokens_sum, batch_seqlens, delta_time):
hidden_size = config.hidden_size
vocab_size = config.vocab_size
Expand Down Expand Up @@ -549,6 +641,8 @@ def _estimate_unknown_flops(config, tokens_sum, batch_seqlens, delta_time):
"qwen3_moe": _estimate_qwen2_moe_flops,
"qwen3_vl": _estimate_qwen3_vl_flops,
"qwen3_vl_moe": _estimate_qwen3_vl_moe_flops,
"qwen3_5": _estimate_qwen3_5_flops,
"qwen3_5_moe": _estimate_qwen3_5_flops,
"deepseek_v3": _estimate_deepseek_v3_flops,
"minicpmv": _estimate_qwen2_flops,
"minicpmo": _estimate_qwen2_flops,
Expand Down