From 2248c6915242be497d74ca37bee91810f5f20ad0 Mon Sep 17 00:00:00 2001 From: zitai-wang <2531131993@qq.com> Date: Wed, 26 Aug 2026 14:35:21 +0800 Subject: [PATCH 1/5] feat(checkpoint): support packed-section TP mapping --- areno/engine/checkpoints/common.py | 83 +++++++++++++++++++++++++++++- 1 file changed, 82 insertions(+), 1 deletion(-) diff --git a/areno/engine/checkpoints/common.py b/areno/engine/checkpoints/common.py index 16358c49..a17a085e 100644 --- a/areno/engine/checkpoints/common.py +++ b/areno/engine/checkpoints/common.py @@ -98,6 +98,16 @@ class MergedColumnSpec: keys: tuple[str, ...] +@dataclass(frozen=True, slots=True) +class PackedSectionColumnSpec: + """One HF tensor whose semantic row sections are TP-sharded separately.""" + + key: str + tensor_attr: str + global_sizes_attr: str + local_sizes_attr: str + + @dataclass(frozen=True, slots=True) class KSharedQKVColumnSpec: """QKV load spec for checkpoints where later layers may share K/V.""" @@ -403,6 +413,8 @@ def save_checkpoint_weights( source_path: str | None, spec: CheckpointSpec, extra_tensors_fn: Callable[[CheckpointTensorStore], None] | None = None, + *, + copy_passthrough: bool = True, ) -> str | None: """Save a tensor-parallel model as a HF sharded safetensors checkpoint.""" @@ -430,7 +442,7 @@ def save_checkpoint_weights( writer.write(tensors, "extra-tensors") tensors.clear() saved_path = writer.finish() - if saved_path is not None and source_path is not None: + if copy_passthrough and saved_path is not None and source_path is not None: copy_source_passthrough_weights( source_path, saved_path, protected_prefix=_protected_prefix_from_top_level(spec.top_level) ) @@ -599,6 +611,9 @@ def load_layer_op( if isinstance(op, MergedColumnSpec): load_merged_column_spec(module, index, prefix, op, rank, world_size) return + if isinstance(op, PackedSectionColumnSpec): + load_packed_section_column_spec(module, index, prefix, op, rank, world_size) + return if isinstance(op, KSharedQKVColumnSpec): load_k_shared_qkv_column_spec(module, index, prefix, op, rank, world_size) return @@ -642,6 +657,9 @@ def save_layer_op( if isinstance(op, SplitColumnSpec): save_split_column_spec(tensors, module, prefix, op) return + if isinstance(op, PackedSectionColumnSpec): + save_packed_section_column_spec(tensors, module, prefix, op) + return if isinstance(op, RangedSplitColumnSpec): save_ranged_split_column_spec(tensors, module, prefix, op) return @@ -751,6 +769,47 @@ def load_merged_column_spec( copy_merged_column_from_index(dst, index, tensor_keys, rank, world_size) +def load_packed_section_column_spec( + module: nn.Module, + index: SafetensorsIndex, + prefix: str, + spec: PackedSectionColumnSpec, + rank: int, + world_size: int, +) -> None: + """Shard each row section of one packed HF tensor independently.""" + + dst = attr_path(module, spec.tensor_attr) + global_sizes = tuple(int(size) for size in attr_path(module, spec.global_sizes_attr)) + local_sizes = tuple(int(size) for size in attr_path(module, spec.local_sizes_attr)) + if len(global_sizes) != len(local_sizes): + raise ValueError("packed-section global and local size counts differ") + ranges = tuple(_shard_range(size, rank, world_size) for size in global_sizes) + expected_local_sizes = tuple(end - start for start, end in ranges) + if local_sizes != expected_local_sizes: + raise ValueError(f"packed-section local sizes {local_sizes} do not match TP shard sizes {expected_local_sizes}") + if dst.shape[0] != sum(local_sizes): + raise ValueError(f"packed-section destination has {dst.shape[0]} rows, expected {sum(local_sizes)}") + + tensor_key = key(spec.key, prefix) + filename = index.weight_map.get(tensor_key) + if filename is None: + raise KeyError(f"missing HF weight {tensor_key}") + with safe_open(index.model_path / filename, framework="pt", device="cpu") as handle: + source = handle.get_slice(tensor_key) + source_shape = tuple(source.get_shape()) + expected_shape = (sum(global_sizes), *dst.shape[1:]) + if source_shape != expected_shape: + raise ValueError(f"checkpoint tensor {tensor_key} has shape {source_shape}, expected {expected_shape}") + source_offset = 0 + destination_offset = 0 + for global_size, local_size, (start, end) in zip(global_sizes, local_sizes, ranges, strict=True): + shard = source[source_offset + start : source_offset + end] + dst[destination_offset : destination_offset + local_size].copy_(shard.to(dtype=dst.dtype)) + source_offset += global_size + destination_offset += local_size + + def load_k_shared_qkv_column_spec( module: nn.Module, index: SafetensorsIndex, prefix: str, spec: KSharedQKVColumnSpec, rank: int, world_size: int ) -> None: @@ -810,6 +869,28 @@ def save_split_column_spec( tensors[key(template, prefix)] = tensor +def save_packed_section_column_spec( + tensors: dict[str, torch.Tensor | None], + module: nn.Module, + prefix: str, + spec: PackedSectionColumnSpec, +) -> None: + """Gather local packed sections into their original single HF tensor.""" + + tensor = attr_path(module, spec.tensor_attr) + global_sizes = tuple(int(size) for size in attr_path(module, spec.global_sizes_attr)) + local_sizes = [int(size) for size in attr_path(module, spec.local_sizes_attr)] + world_size = get_tp_context().world_size + if any( + global_size != local_size * world_size + for global_size, local_size in zip(global_sizes, local_sizes, strict=True) + ): + raise ValueError("packed-section sizes are incompatible with the tensor-parallel world size") + if tensor.shape[0] != sum(local_sizes): + raise ValueError(f"packed-section source has {tensor.shape[0]} rows, expected {sum(local_sizes)}") + tensors[key(spec.key, prefix)] = gather_tensor_parallel_split_column_tensor(tensor, local_sizes) + + def save_ranged_split_column_spec( tensors: dict[str, torch.Tensor | None], module: nn.Module, prefix: str, spec: RangedSplitColumnSpec ) -> None: From 20ef04900975082d7cf522c31f1367dee088979a Mon Sep 17 00:00:00 2001 From: zitai-wang <2531131993@qq.com> Date: Wed, 26 Aug 2026 14:36:24 +0800 Subject: [PATCH 2/5] feat(models): add Phi-4 text-only adapter --- areno/engine/layers/attention.py | 23 +- areno/engine/runtime/decode_graph.py | 13 + areno/models/__init__.py | 8 + areno/models/phi4mm/__init__.py | 21 ++ areno/models/phi4mm/checkpoint.py | 162 ++++++++++ areno/models/phi4mm/model.py | 467 +++++++++++++++++++++++++++ 6 files changed, 691 insertions(+), 3 deletions(-) create mode 100644 areno/models/phi4mm/__init__.py create mode 100644 areno/models/phi4mm/checkpoint.py create mode 100644 areno/models/phi4mm/model.py diff --git a/areno/engine/layers/attention.py b/areno/engine/layers/attention.py index 0f1b05bb..20ef88cb 100644 --- a/areno/engine/layers/attention.py +++ b/areno/engine/layers/attention.py @@ -32,7 +32,7 @@ class CausalSelfAttention(nn.Module): reduces across ranks to reassemble the full hidden state. """ - def __init__(self, config: ModelConfig, layer_idx: int): + def __init__(self, config: ModelConfig, layer_idx: int, *, rotary_embedding: nn.Module | None = None): super().__init__() ctx = get_tp_context() self.layer_idx = layer_idx @@ -58,7 +58,11 @@ def __init__(self, config: ModelConfig, layer_idx: int): # Row-parallel output projection: input is already sharded along # head dimension, output is all-reduced across ranks. self.o_proj = RowParallelLinear(self.num_heads * self.head_dim, config.hidden_size, bias=False) - self.rope = RotaryEmbedding(config.head_dim, config.max_position_embeddings, config.rope_theta) + self.rope = ( + rotary_embedding + if rotary_embedding is not None + else RotaryEmbedding(config.head_dim, config.max_position_embeddings, config.rope_theta) + ) # Optional per-head QK normalization (used by some recent models). self.q_norm = RMSNorm(config.head_dim, config.rms_norm_eps) if config.qk_norm else None self.k_norm = RMSNorm(config.head_dim, config.rms_norm_eps) if config.qk_norm else None @@ -93,7 +97,7 @@ def forward( k = self.k_norm(k) # Rotary embedding is applied on the head dim using position-indexed # cos/sin tables; positions are broadcast across heads. - q, k = self.rope(q, k, position_ids) + q, k = self.apply_rotary(q, k, position_ids, train_meta, infer_meta) # Presence of infer_meta selects the paged KV-cache backend; otherwise # we run the training-mode FlashAttention (padded or varlen packed). @@ -101,6 +105,19 @@ def forward( return self.forward_infer(q, k, v, infer_meta) return self.forward_train(q, k, v, train_meta) + def apply_rotary( + self, + q: torch.Tensor, + k: torch.Tensor, + position_ids: torch.Tensor, + train_meta: TrainMeta | None, + infer_meta: InferMeta | None, + ) -> tuple[torch.Tensor, torch.Tensor]: + """Apply the model's rotary embedding, with a model override hook.""" + + del train_meta, infer_meta + return self.rope(q, k, position_ids) + def forward_train( self, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, train_meta: TrainMeta | None ) -> torch.Tensor: diff --git a/areno/engine/runtime/decode_graph.py b/areno/engine/runtime/decode_graph.py index 79f0d0eb..339e29d4 100644 --- a/areno/engine/runtime/decode_graph.py +++ b/areno/engine/runtime/decode_graph.py @@ -81,6 +81,7 @@ def __init__( """Allocate static input buffers and the `InferMeta` baked into capture.""" self.model = model + self.decode_cache_length_limit = getattr(model, "decode_cache_length_limit", None) self.bucket = bucket self.scratch_block = scratch_block self.scratch_recurrent_slot = scratch_recurrent_slot @@ -156,6 +157,7 @@ def replay_tensors( actual = int(input_ids.numel()) if actual > self.bucket: raise ValueError(f"decode payload has {actual} tokens, graph bucket is {self.bucket}") + _validate_decode_cache_length(cache_seqlens, actual, self.decode_cache_length_limit) # Copy the live values into the captured-stable buffers. The graph # was recorded against these buffer addresses so `copy_` here is what @@ -186,3 +188,14 @@ def replay_tensors( self.graph.replay() assert self.logits_shard is not None return self.logits_shard + + +def _validate_decode_cache_length( + cache_seqlens: torch.Tensor, + actual: int, + limit: int | None, +) -> None: + if limit is not None and actual and int(cache_seqlens[:actual].max().item()) >= limit: + raise ValueError( + "cached decode cannot cross the model's rotary-factor boundary; run a full long-context prefill" + ) diff --git a/areno/models/__init__.py b/areno/models/__init__.py index 83540786..cb74c51e 100644 --- a/areno/models/__init__.py +++ b/areno/models/__init__.py @@ -25,6 +25,13 @@ def _register_qwen35() -> None: register_adapter(Qwen35Adapter()) +def _register_phi4mm() -> None: + from areno.models.phi4mm import Phi4MMAdapter + from areno.models.registry import register_adapter + + register_adapter(Phi4MMAdapter()) + + def _register_bailing() -> None: from areno.models.bailing import BailingMoeLinearV2Adapter from areno.models.registry import register_adapter @@ -71,6 +78,7 @@ def _register_olmo2() -> None: "llama": _register_llama, "qwen3": _register_qwen3, "qwen3_5": _register_qwen35, + "phi4mm": _register_phi4mm, "bailing": _register_bailing, "bailing_v3": _register_bailing_v3, "gemma4": _register_gemma4, diff --git a/areno/models/phi4mm/__init__.py b/areno/models/phi4mm/__init__.py new file mode 100644 index 00000000..3cf01600 --- /dev/null +++ b/areno/models/phi4mm/__init__.py @@ -0,0 +1,21 @@ +"""Phi-4-Multimodal language-backbone adapter.""" + +from __future__ import annotations + +from areno.models.phi4mm.model import ( + Phi4MMAdapter, + Phi4MMAttention, + Phi4MMDecoderLayer, + Phi4MMForCausalLM, + Phi4MMLongRoPEScaledRotaryEmbedding, + Phi4MMModel, +) + +__all__ = [ + "Phi4MMAdapter", + "Phi4MMAttention", + "Phi4MMDecoderLayer", + "Phi4MMForCausalLM", + "Phi4MMLongRoPEScaledRotaryEmbedding", + "Phi4MMModel", +] diff --git a/areno/models/phi4mm/checkpoint.py b/areno/models/phi4mm/checkpoint.py new file mode 100644 index 00000000..9fcd4d9a --- /dev/null +++ b/areno/models/phi4mm/checkpoint.py @@ -0,0 +1,162 @@ +"""Strict text-only checkpoint mapping for Phi-4-Multimodal.""" + +from __future__ import annotations + +import re +from dataclasses import dataclass +from pathlib import Path + +from torch import nn + +from areno.engine.checkpoints.common import ( + CheckpointSpec, + LayerSpec, + PackedSectionColumnSpec, + ParallelTensorSpec, + ReplicatedTensorSpec, + TopLevelSpec, + load_checkpoint_weights, + save_checkpoint_weights, +) +from areno.engine.checkpoints.io import SafetensorsIndex +from areno.engine.parallel.context import get_tp_context + +TOP_LEVEL_SPEC = TopLevelSpec( + embedding_key="model.embed_tokens.weight", + embedding_attr="model.embed_tokens", + norm_key="model.norm.weight", + norm_attr="model.norm.weight", +) +LAYER_NORM_SPECS = ( + ReplicatedTensorSpec("{prefix}.input_layernorm.weight", "input_layernorm.weight"), + ReplicatedTensorSpec("{prefix}.post_attention_layernorm.weight", "post_attention_layernorm.weight"), +) +QKV_SPEC = PackedSectionColumnSpec( + key="{prefix}.self_attn.qkv_proj.base_layer.weight", + tensor_attr="self_attn.qkv_proj.weight", + global_sizes_attr="self_attn.qkv_proj.out_features", + local_sizes_attr="self_attn.qkv_proj.local_out_features", +) +ATTN_OUT_SPEC = ParallelTensorSpec( + "{prefix}.self_attn.o_proj.base_layer.weight", + "self_attn.o_proj.weight", + 1, +) +GATE_UP_SPEC = PackedSectionColumnSpec( + key="{prefix}.mlp.gate_up_proj.base_layer.weight", + tensor_attr="mlp.gate_up_proj.weight", + global_sizes_attr="mlp.gate_up_proj.out_features", + local_sizes_attr="mlp.gate_up_proj.local_out_features", +) +MLP_DOWN_SPEC = ParallelTensorSpec( + "{prefix}.mlp.down_proj.base_layer.weight", + "mlp.down_proj.weight", + 1, +) +LAYER_SPEC = LayerSpec( + prefix="model.layers.{layer}", + replicated=LAYER_NORM_SPECS, + load_ops=(QKV_SPEC, ATTN_OUT_SPEC, GATE_UP_SPEC, MLP_DOWN_SPEC), + save_ops=(QKV_SPEC, ATTN_OUT_SPEC, GATE_UP_SPEC, MLP_DOWN_SPEC), +) +CHECKPOINT_SPEC = CheckpointSpec(top_level=TOP_LEVEL_SPEC, layer=LAYER_SPEC) + +_LAYER_BASE_SUFFIXES = ( + "input_layernorm.weight", + "post_attention_layernorm.weight", + "self_attn.qkv_proj.base_layer.weight", + "self_attn.o_proj.base_layer.weight", + "mlp.gate_up_proj.base_layer.weight", + "mlp.down_proj.base_layer.weight", +) +_LORA_PATTERN = re.compile( + r"^model\.layers\.(\d+)\." + r"(?:self_attn\.(?:qkv_proj|o_proj)|mlp\.(?:gate_up_proj|down_proj))\." + r"lora_[AB]\.(vision|speech)\.weight$" +) + + +@dataclass(frozen=True, slots=True) +class Phi4MMCheckpointAudit: + total: int + consumed: int + vision_lora_skipped: int + speech_lora_skipped: int + vision_skipped: int + audio_skipped: int + unknown: int + + +def _required_base_keys(num_hidden_layers: int) -> set[str]: + required = {"model.embed_tokens.weight", "model.norm.weight"} + for layer in range(num_hidden_layers): + required.update(f"model.layers.{layer}.{suffix}" for suffix in _LAYER_BASE_SUFFIXES) + return required + + +def audit_phi4mm_checkpoint(model_path: str | Path, num_hidden_layers: int) -> Phi4MMCheckpointAudit: + """Classify every checkpoint key and reject missing or unknown tensors.""" + + index = SafetensorsIndex(model_path, progress=False) + try: + checkpoint_keys = set(index.weight_map) + finally: + index.close() + required = _required_base_keys(num_hidden_layers) + missing = sorted(required - checkpoint_keys) + if missing: + preview = ", ".join(missing[:5]) + raise ValueError(f"Phi4MM checkpoint is missing {len(missing)} required base-language tensors: {preview}") + + counts = {"vision_lora": 0, "speech_lora": 0, "vision": 0, "audio": 0} + unknown = [] + for tensor_key in checkpoint_keys - required: + lora_match = _LORA_PATTERN.fullmatch(tensor_key) + if lora_match is not None and int(lora_match.group(1)) < num_hidden_layers: + counts[f"{lora_match.group(2)}_lora"] += 1 + elif tensor_key.startswith("model.embed_tokens_extend.image_embed."): + counts["vision"] += 1 + elif tensor_key.startswith("model.embed_tokens_extend.audio_embed."): + counts["audio"] += 1 + else: + unknown.append(tensor_key) + if unknown: + preview = ", ".join(sorted(unknown)[:5]) + raise ValueError(f"Phi4MM checkpoint contains {len(unknown)} unknown tensors: {preview}") + return Phi4MMCheckpointAudit( + total=len(checkpoint_keys), + consumed=len(required), + vision_lora_skipped=counts["vision_lora"], + speech_lora_skipped=counts["speech_lora"], + vision_skipped=counts["vision"], + audio_skipped=counts["audio"], + unknown=0, + ) + + +def load_phi4mm_weights(model: nn.Module, model_path: str | Path) -> Phi4MMCheckpointAudit: + """Audit and load the supported Phi-4 base-language tensors.""" + + model.config.validate_tp(get_tp_context().world_size) + audit = audit_phi4mm_checkpoint(model_path, len(model.layers)) + load_checkpoint_weights(model, str(model_path), CHECKPOINT_SPEC) + if model.lm_head.weight is not model.model.embed_tokens.weight: + raise RuntimeError("Phi4MM embedding and LM head weight tying was lost during checkpoint loading") + return audit + + +def save_phi4mm_weights( + model: nn.Module, + output_path: str | Path, + source_path: str | Path | None, +) -> str | None: + """Save only Phi-4 base-language weights in official HF key layout.""" + + model.config.validate_tp(get_tp_context().world_size) + return save_checkpoint_weights( + model, + str(output_path), + None if source_path is None else str(source_path), + CHECKPOINT_SPEC, + copy_passthrough=False, + ) diff --git a/areno/models/phi4mm/model.py b/areno/models/phi4mm/model.py new file mode 100644 index 00000000..45118cf3 --- /dev/null +++ b/areno/models/phi4mm/model.py @@ -0,0 +1,467 @@ +"""Phi-4-Multimodal language-backbone adapter. + +PR1 intentionally supports the checkpoint's text path only. The vision and +audio towers and their modality-specific LoRA adapters are not runtime model +components here. +""" + +from __future__ import annotations + +import math +from pathlib import Path +from typing import Any + +import torch +from torch import nn + +from areno.accel.ops import is_cuda_graph_capturing +from areno.engine.config import ModelConfig, _parse_dtype +from areno.engine.layers.attention import CausalSelfAttention +from areno.engine.layers.mlp import GatedMLP +from areno.engine.layers.norm import RMSNorm +from areno.engine.layers.vocab import VocabParallelEmbedding, VocabParallelLMHead +from areno.engine.parallel.collectives import ( + scatter_to_sequence_parallel_region, + sequence_parallel_region, +) +from areno.engine.runtime.metadata import InferMeta, TrainMeta +from areno.engine.runtime.recompute import checkpoint_layer +from areno.models.base import CausalLMOutput, ModelAdapter + + +def _require_bool(hf_config: dict[str, Any], key: str, expected: bool) -> None: + value = bool(hf_config.get(key, expected)) + if value is not expected: + raise ValueError(f"Phi4MM requires {key}={expected}, got {value}") + + +def _validated_longrope(hf_config: dict[str, Any], rotary_dim: int) -> dict[str, Any]: + rope = hf_config.get("rope_scaling") + if not isinstance(rope, dict): + raise ValueError("Phi4MM requires a rope_scaling mapping") + if set(rope) != {"type", "short_factor", "long_factor"}: + raise ValueError("Phi4MM rope_scaling must contain exactly: type, short_factor, long_factor") + if rope["type"] != "longrope": + raise ValueError(f"Phi4MM only supports rope_scaling.type='longrope', got {rope['type']!r}") + + expected_factors = rotary_dim // 2 + normalized = {"type": "longrope"} + for key in ("short_factor", "long_factor"): + factors = rope[key] + if not isinstance(factors, list) or len(factors) != expected_factors: + raise ValueError(f"Phi4MM {key} must contain {expected_factors} values") + if any(isinstance(value, bool) or not isinstance(value, (int, float)) or value <= 0 for value in factors): + raise ValueError(f"Phi4MM {key} values must be positive numbers") + normalized[key] = tuple(float(value) for value in factors) + return normalized + + +def _rotate_half(x: torch.Tensor) -> torch.Tensor: + first, second = x.chunk(2, dim=-1) + return torch.cat((-second, first), dim=-1) + + +class Phi4MMLongRoPEScaledRotaryEmbedding(nn.Module): + """Official Phi-4 partial LongRoPE math without per-layer position caches.""" + + def __init__(self, config: ModelConfig): + super().__init__() + if config.hf_text_config is None: + raise ValueError("Phi4MM requires the validated HF text config") + self.dim = int(config.head_dim * config.partial_rotary_factor) + if self.dim <= 0 or self.dim % 2: + raise ValueError("Phi4MM rotary dimension must be a positive even number") + rope_scaling = config.hf_text_config["rope_scaling"] + expected_factors = self.dim // 2 + short_factor = rope_scaling["short_factor"] + long_factor = rope_scaling["long_factor"] + if len(short_factor) != expected_factors or len(long_factor) != expected_factors: + raise ValueError(f"Phi4MM short_factor and long_factor must contain {expected_factors} values") + + self.max_position_embeddings = int(config.max_position_embeddings) + self.original_max_position_embeddings = int(config.hf_text_config["original_max_position_embeddings"]) + inv_freq_shape = torch.arange(0, self.dim, 2, dtype=torch.int64).float() / self.dim + base_freq = config.rope_theta**inv_freq_shape + self.register_buffer( + "short_inv_freq", 1.0 / (torch.tensor(short_factor, dtype=torch.float32) * base_freq), persistent=False + ) + self.register_buffer( + "long_inv_freq", 1.0 / (torch.tensor(long_factor, dtype=torch.float32) * base_freq), persistent=False + ) + scale = self.max_position_embeddings / self.original_max_position_embeddings + self.scaling_factor = ( + 1.0 if scale <= 1.0 else math.sqrt(1.0 + math.log(scale) / math.log(self.original_max_position_embeddings)) + ) + + def _apply(self, fn): + super()._apply(fn) + # Long-context phases must remain FP32 even when model weights are cast. + self.short_inv_freq = self.short_inv_freq.float() + self.long_inv_freq = self.long_inv_freq.float() + return self + + @torch.no_grad() + def cos_sin( + self, + x: torch.Tensor, + position_ids: torch.Tensor, + sequence_length: int | None = None, + ) -> tuple[torch.Tensor, torch.Tensor]: + if sequence_length is None: + sequence_length = int(torch.max(position_ids).item()) + 1 + inv_freq = ( + self.long_inv_freq if sequence_length > self.original_max_position_embeddings else self.short_inv_freq + ) + expanded_inv_freq = inv_freq[None, :, None].float().expand(position_ids.shape[0], -1, 1) + expanded_positions = position_ids[:, None, :].float() + device_type = x.device.type if x.device.type != "mps" else "cpu" + with torch.autocast(device_type=device_type, enabled=False): + freqs = (expanded_inv_freq @ expanded_positions).transpose(1, 2) + embedding = torch.cat((freqs, freqs), dim=-1) + cos = embedding.cos() * self.scaling_factor + sin = embedding.sin() * self.scaling_factor + return cos.to(dtype=x.dtype), sin.to(dtype=x.dtype) + + def forward( + self, + q: torch.Tensor, + k: torch.Tensor, + position_ids: torch.Tensor, + sequence_length: int | None = None, + ) -> tuple[torch.Tensor, torch.Tensor]: + cos, sin = self.cos_sin(q, position_ids, sequence_length) + cos = cos.unsqueeze(2) + sin = sin.unsqueeze(2) + q_rot, q_pass = q[..., : self.dim], q[..., self.dim :] + k_rot, k_pass = k[..., : self.dim], k[..., self.dim :] + q_embed = torch.cat((q_rot * cos + _rotate_half(q_rot) * sin, q_pass), dim=-1) + k_embed = torch.cat((k_rot * cos + _rotate_half(k_rot) * sin, k_pass), dim=-1) + return q_embed, k_embed + + +def _phi4mm_longrope_sequence_length( + position_ids: torch.Tensor, + train_meta: TrainMeta | None, + infer_meta: InferMeta | None, + original_max_position_embeddings: int, +) -> int: + if infer_meta is not None and infer_meta.mode == "decode": + if infer_meta.cache_seqlens is None: + raise ValueError("Phi4MM decode requires cache_seqlens for LongRoPE selection") + sequence_length = int(infer_meta.cache_seqlens.max().item()) + 1 + if sequence_length > original_max_position_embeddings: + raise ValueError( + "Phi4MM cached decode cannot cross the LongRoPE boundary because cached keys may use short factors; " + "run a full long-context prefill" + ) + return sequence_length + + if infer_meta is not None: + sequence_length = int(position_ids.max().item()) + 1 + if sequence_length > original_max_position_embeddings: + if infer_meta.cu_seqlens is None: + raise ValueError("Phi4MM prefill requires cu_seqlens for LongRoPE boundary validation") + starts = infer_meta.cu_seqlens[:-1].to(dtype=torch.long) + flat_positions = position_ids.reshape(-1) + if bool(torch.any(flat_positions[starts] != 0)): + raise ValueError( + "Phi4MM chunked prefill cannot cross the LongRoPE boundary because cached keys use short factors; " + "increase the prefill token budget and run a full prefill" + ) + return sequence_length + + if train_meta is not None and train_meta.max_seqlen is not None: + return int(train_meta.max_seqlen) + return int(position_ids.shape[-1]) + + +class Phi4MMAttention(CausalSelfAttention): + """AReno GQA attention with a Phi-owned rotary implementation.""" + + def __init__(self, config: ModelConfig, layer_idx: int): + if config.qk_norm: + raise ValueError("Phi4MMAttention requires qk_norm=False") + super().__init__(config, layer_idx, rotary_embedding=Phi4MMLongRoPEScaledRotaryEmbedding(config)) + + def apply_rotary( + self, + q: torch.Tensor, + k: torch.Tensor, + position_ids: torch.Tensor, + train_meta: TrainMeta | None, + infer_meta: InferMeta | None, + ) -> tuple[torch.Tensor, torch.Tensor]: + if infer_meta is not None and infer_meta.mode == "decode" and is_cuda_graph_capturing(q): + # DecodeGraph validates its dynamic cache lengths before replay. + # Capture itself always records the supported short-factor path. + sequence_length = self.rope.original_max_position_embeddings + else: + sequence_length = _phi4mm_longrope_sequence_length( + position_ids, + train_meta, + infer_meta, + self.rope.original_max_position_embeddings, + ) + return self.rope(q, k, position_ids, sequence_length) + + +class Phi4MMDecoderLayer(nn.Module): + """Phi-4 pre-norm decoder block composed from AReno shared layers.""" + + def __init__(self, config: ModelConfig, layer_idx: int): + super().__init__() + self.input_layernorm = RMSNorm(config.hidden_size, config.rms_norm_eps) + self.self_attn = Phi4MMAttention(config, layer_idx) + self.post_attention_layernorm = RMSNorm(config.hidden_size, config.rms_norm_eps) + self.mlp = GatedMLP(config) + + def forward( + self, + hidden_states: torch.Tensor, + position_ids: torch.Tensor, + train_meta: TrainMeta | None = None, + infer_meta: InferMeta | None = None, + ) -> torch.Tensor: + residual = hidden_states + hidden_states = self.input_layernorm(hidden_states) + hidden_states = residual + self.self_attn(hidden_states, position_ids, train_meta, infer_meta) + residual = hidden_states + hidden_states = self.post_attention_layernorm(hidden_states) + return residual + self.mlp(hidden_states) + + +class Phi4MMModel(nn.Module): + """Text-only Phi-4 transformer body.""" + + def __init__(self, config: ModelConfig): + super().__init__() + self.config = config + self.embed_tokens = VocabParallelEmbedding(config.vocab_size, config.hidden_size, dtype=config.dtype) + self.layers = nn.ModuleList([Phi4MMDecoderLayer(config, index) for index in range(config.num_hidden_layers)]) + self.norm = RMSNorm(config.hidden_size, config.rms_norm_eps) + + def forward( + self, + input_ids: torch.Tensor, + position_ids: torch.Tensor | None = None, + train_meta: TrainMeta | None = None, + infer_meta: InferMeta | None = None, + ) -> torch.Tensor: + if position_ids is None: + position_ids = torch.arange(input_ids.shape[1], device=input_ids.device).unsqueeze(0).expand_as(input_ids) + hidden_states = self.embed_tokens(input_ids) + use_sequence_parallel = bool(train_meta is not None and train_meta.sequence_parallel) + if use_sequence_parallel: + hidden_states = scatter_to_sequence_parallel_region(hidden_states) + with sequence_parallel_region(use_sequence_parallel): + for layer in self.layers: + hidden_states = checkpoint_layer( + layer, + hidden_states, + position_ids, + train_meta, + infer_meta, + train_meta=train_meta, + infer_meta=infer_meta, + ) + return self.norm(hidden_states) + + +class Phi4MMForCausalLM(nn.Module): + """Text-only Phi-4 causal LM with a truly tied vocab-parallel head.""" + + def __init__(self, config: ModelConfig): + super().__init__() + if not config.tie_word_embeddings: + raise ValueError("Phi4MMForCausalLM requires tied word embeddings") + self.config = config + self.decode_cache_length_limit = int(config.hf_text_config["original_max_position_embeddings"]) + self.model = Phi4MMModel(config) + self.lm_head = VocabParallelLMHead(config.hidden_size, config.vocab_size, dtype=config.dtype) + self._tie_word_embeddings() + + def _tie_word_embeddings(self) -> None: + embedding = self.model.embed_tokens + if (self.lm_head.vocab_start, self.lm_head.vocab_end) != (embedding.vocab_start, embedding.vocab_end): + raise ValueError("Phi4MM embedding and LM head use different TP vocabulary ranges") + if self.lm_head.weight.shape != embedding.weight.shape: + raise ValueError("Phi4MM embedding and LM head local weight shapes differ") + self.lm_head.weight = embedding.weight + + @property + def layers(self) -> nn.ModuleList: + """Expose decoder layers to the shared checkpoint machinery.""" + + return self.model.layers + + def forward( + self, + input_ids: torch.Tensor, + position_ids: torch.Tensor | None = None, + train_meta: TrainMeta | None = None, + infer_meta: InferMeta | None = None, + ) -> CausalLMOutput: + use_sequence_parallel = bool(train_meta is not None and train_meta.sequence_parallel) + with sequence_parallel_region(use_sequence_parallel): + hidden_states = self.model(input_ids, position_ids, train_meta, infer_meta) + logits_shard = self.lm_head(hidden_states) + return CausalLMOutput(logits_shard=logits_shard, hidden_states=hidden_states) + + def set_kv_caches( + self, kv_caches: list[tuple[torch.Tensor, torch.Tensor]], *, num_slots: int | None = None + ) -> None: + """Bind one paged KV-cache pair to each decoder layer.""" + del num_slots + if len(kv_caches) != len(self.layers): + raise ValueError(f"expected {len(self.layers)} layer caches, got {len(kv_caches)}") + for layer, (k_cache, v_cache) in zip(self.layers, kv_caches, strict=True): + layer.self_attn.set_kv_cache(k_cache, v_cache) + + @torch.no_grad() + def prepare_infer_weights(self) -> None: + return None + + @torch.no_grad() + def clear_infer_weights(self) -> None: + return None + + @torch.no_grad() + def offload_train_weights(self) -> None: + return None + + @torch.no_grad() + def onload_train_weights(self, device: torch.device) -> None: + del device + return None + + @torch.no_grad() + def finalize_router_expert_bias(self, tp_group, dp_group) -> None: + del tp_group, dp_group + return None + + def allocate_kv_caches( + self, num_blocks: int, block_size: int, device: torch.device + ) -> list[tuple[torch.Tensor, torch.Tensor]]: + """Allocate the standard paged GQA cache layout for every layer.""" + caches = [] + for layer in self.layers: + attention = layer.self_attn + shape = (num_blocks, block_size, attention.local_kv_heads, attention.head_dim) + caches.append( + ( + torch.empty(shape, device=device, dtype=self.config.dtype), + torch.empty(shape, device=device, dtype=self.config.dtype), + ) + ) + return caches + + def clear_kv_caches(self) -> None: + for layer in self.layers: + layer.self_attn.clear_kv_cache() + + @torch.no_grad() + def reset_kv_caches(self) -> None: + return None + + @torch.no_grad() + def offload_kv_caches(self) -> None: + for layer in self.layers: + attention = layer.self_attn + if attention.k_cache.numel() > 0: + attention.k_cache = attention.k_cache.to(device="cpu") + if attention.v_cache.numel() > 0: + attention.v_cache = attention.v_cache.to(device="cpu") + attention.infer_backend = None + + @torch.no_grad() + def onload_kv_caches(self, device: torch.device) -> bool: + found = False + for layer in self.layers: + attention = layer.self_attn + if attention.k_cache.numel() > 0: + found = True + if attention.k_cache.device != device: + attention.k_cache = attention.k_cache.to(device=device) + if attention.v_cache.numel() > 0 and attention.v_cache.device != device: + attention.v_cache = attention.v_cache.to(device=device) + return found + + +class Phi4MMAdapter(ModelAdapter): + """Translate the official Phi-4-Multimodal config into AReno semantics.""" + + name = "phi4mm" + + def match_hf_config(self, hf_config: dict[str, Any]) -> bool: + return str(hf_config.get("model_type", "")).lower() == self.name + + def config_from_hf(self, hf_config: dict[str, Any]) -> ModelConfig: + hidden_size = int(hf_config["hidden_size"]) + num_attention_heads = int(hf_config["num_attention_heads"]) + if hidden_size % num_attention_heads != 0: + raise ValueError("Phi4MM hidden_size must be divisible by num_attention_heads") + head_dim = hidden_size // num_attention_heads + partial_rotary_factor = float(hf_config.get("partial_rotary_factor", 1.0)) + if not 0.0 < partial_rotary_factor <= 1.0: + raise ValueError("Phi4MM partial_rotary_factor must be in (0, 1]") + rotary_dim = int(head_dim * partial_rotary_factor) + if rotary_dim <= 0 or rotary_dim % 2 != 0: + raise ValueError("Phi4MM rotary dimension must be a positive even number") + + if str(hf_config.get("hidden_act", "silu")) != "silu": + raise ValueError("Phi4MM language backbone requires hidden_act='silu'") + _require_bool(hf_config, "attention_bias", False) + _require_bool(hf_config, "mlp_bias", False) + _require_bool(hf_config, "lm_head_bias", False) + _require_bool(hf_config, "tie_word_embeddings", True) + + original_max_position_embeddings = int(hf_config.get("original_max_position_embeddings", 4096)) + max_position_embeddings = int(hf_config.get("max_position_embeddings", original_max_position_embeddings)) + if original_max_position_embeddings <= 0 or max_position_embeddings < original_max_position_embeddings: + raise ValueError("Phi4MM max_position_embeddings must be at least original_max_position_embeddings > 0") + rope_scaling = _validated_longrope(hf_config, rotary_dim) + + # Preserve the validated LongRoPE fields for the Phi-specific rotary implementation. + text_config = dict(hf_config) + text_config["rope_scaling"] = rope_scaling + text_config["original_max_position_embeddings"] = original_max_position_embeddings + + return ModelConfig( + model_type=self.name, + checkpoint_prefix="model", + vocab_size=int(hf_config["vocab_size"]), + pad_token_id=int(hf_config.get("pad_token_id", 0) or 0), + hidden_size=hidden_size, + intermediate_size=int(hf_config["intermediate_size"]), + num_hidden_layers=int(hf_config["num_hidden_layers"]), + num_attention_heads=num_attention_heads, + num_key_value_heads=int(hf_config.get("num_key_value_heads", num_attention_heads)), + head_dim=head_dim, + rms_norm_eps=float(hf_config.get("rms_norm_eps", 1e-5)), + rope_theta=float(hf_config.get("rope_theta", 10_000.0)), + max_position_embeddings=max_position_embeddings, + tie_word_embeddings=True, + qkv_bias=False, + qk_norm=False, + dtype=_parse_dtype(hf_config.get("torch_dtype") or hf_config.get("dtype")), + hidden_act="silu", + sliding_window=hf_config.get("sliding_window"), + partial_rotary_factor=partial_rotary_factor, + sequence_parallel=bool(hf_config.get("sequence_parallel", True)), + hf_text_config=text_config, + ) + + def build(self, config: ModelConfig) -> nn.Module: + if config.model_type != self.name: + raise ValueError(f"Phi4MMAdapter cannot build model_type={config.model_type!r}") + return Phi4MMForCausalLM(config) + + def load_weights(self, model: nn.Module, model_path: str | Path) -> None: + from areno.models.phi4mm.checkpoint import load_phi4mm_weights + + load_phi4mm_weights(model, model_path) + + def save_weights(self, model: nn.Module, output_path: str | Path, source_path: str | Path | None) -> str | None: + from areno.models.phi4mm.checkpoint import save_phi4mm_weights + + return save_phi4mm_weights(model, output_path, source_path) From 0c13a01b41c17782a768bc3601cc2dfca8f2eb59 Mon Sep 17 00:00:00 2001 From: zitai-wang <2531131993@qq.com> Date: Wed, 26 Aug 2026 14:36:31 +0800 Subject: [PATCH 3/5] test(models): add Phi-4 adapter and checkpoint coverage --- tests/test_phi4mm_adapter_cpu.py | 434 ++++++++++++++++++++++++++++ tests/test_phi4mm_checkpoint_cpu.py | 275 ++++++++++++++++++ 2 files changed, 709 insertions(+) create mode 100644 tests/test_phi4mm_adapter_cpu.py create mode 100644 tests/test_phi4mm_checkpoint_cpu.py diff --git a/tests/test_phi4mm_adapter_cpu.py b/tests/test_phi4mm_adapter_cpu.py new file mode 100644 index 00000000..f4c6fcc5 --- /dev/null +++ b/tests/test_phi4mm_adapter_cpu.py @@ -0,0 +1,434 @@ +from __future__ import annotations + +import json +import math + +import pytest +import torch +import torch.nn.functional as F + +import areno.models +from areno.engine.config import ModelConfig, OptimizerConfig +from areno.engine.layers import mlp, norm, vocab +from areno.engine.modeling import build_optimizer +from areno.engine.parallel.collectives import is_sequence_parallel_active +from areno.engine.parallel.context import TPContext, get_tp_context, set_tp_context +from areno.engine.runtime.decode_graph import _validate_decode_cache_length +from areno.engine.runtime.metadata import InferMeta, TrainMeta +from areno.models import registry +from areno.models.phi4mm import Phi4MMAdapter, Phi4MMForCausalLM +from areno.models.phi4mm.model import ( + Phi4MMLongRoPEScaledRotaryEmbedding, + _phi4mm_longrope_sequence_length, +) + + +@pytest.fixture(autouse=True) +def _isolate_tp_context(): + previous_context = get_tp_context() + set_tp_context(TPContext(rank=0, world_size=1, device=torch.device("cpu"), group=None)) + try: + yield + finally: + set_tp_context(previous_context) + + +def _phi4mm_config() -> dict: + return { + "model_type": "phi4mm", + "vocab_size": 200064, + "hidden_size": 3072, + "intermediate_size": 8192, + "num_hidden_layers": 32, + "num_attention_heads": 24, + "num_key_value_heads": 8, + "rms_norm_eps": 1e-5, + "rope_theta": 10_000.0, + "max_position_embeddings": 131072, + "original_max_position_embeddings": 4096, + "partial_rotary_factor": 0.75, + "rope_scaling": { + "type": "longrope", + "short_factor": [1.0] * 48, + "long_factor": [float(index + 1) for index in range(48)], + }, + "sliding_window": 262144, + "hidden_act": "silu", + "attention_bias": False, + "mlp_bias": False, + "lm_head_bias": False, + "tie_word_embeddings": True, + "pad_token_id": 199999, + "torch_dtype": "bfloat16", + } + + +def _tiny_model_config() -> ModelConfig: + return ModelConfig( + model_type="phi4mm", + vocab_size=32, + hidden_size=64, + intermediate_size=128, + num_hidden_layers=2, + num_attention_heads=8, + num_key_value_heads=4, + head_dim=8, + rms_norm_eps=1e-5, + rope_theta=10_000.0, + max_position_embeddings=64, + tie_word_embeddings=True, + qkv_bias=False, + qk_norm=False, + dtype=torch.float32, + hidden_act="silu", + partial_rotary_factor=0.75, + sequence_parallel=False, + attn_backend="native", + hf_text_config={ + "original_max_position_embeddings": 32, + "rope_scaling": { + "type": "longrope", + "short_factor": (1.0, 1.0, 1.0), + "long_factor": (1.0, 2.0, 3.0), + }, + }, + ) + + +@pytest.fixture +def cpu_reference_kernels(monkeypatch): + def embedding(input_ids, weight, vocab_start, vocab_end): + local_ids = input_ids - vocab_start + local_mask = (input_ids >= vocab_start) & (input_ids < vocab_end) + safe_ids = local_ids.masked_fill(~local_mask, 0) + return F.embedding(safe_ids, weight) * local_mask.unsqueeze(-1) + + def rms_norm(x, weight, eps): + normalized = x.float() * torch.rsqrt(x.float().square().mean(dim=-1, keepdim=True) + eps) + return normalized.to(dtype=x.dtype) * weight.to(dtype=x.dtype) + + def silu_and_mul(x): + gate, up = x.chunk(2, dim=-1) + return F.silu(gate) * up + + monkeypatch.setattr(vocab, "areno_vocab_embedding", embedding) + monkeypatch.setattr(norm, "_areno_rmsnorm_no_compile", rms_norm) + monkeypatch.setattr(mlp, "_areno_silu_and_mul_no_compile", silu_and_mul) + + +def test_phi4mm_config_translation_matches_official_language_backbone(): + config = Phi4MMAdapter().config_from_hf(_phi4mm_config()) + + assert config.model_type == "phi4mm" + assert config.vocab_size == 200064 + assert config.hidden_size == 3072 + assert config.intermediate_size == 8192 + assert config.num_hidden_layers == 32 + assert config.num_attention_heads == 24 + assert config.num_key_value_heads == 8 + assert config.head_dim == 128 + assert config.partial_rotary_factor == 0.75 + assert config.qk_norm is False + assert config.qkv_bias is False + assert config.tie_word_embeddings is True + assert config.dtype == torch.bfloat16 + assert config.hf_text_config is not None + assert config.hf_text_config["original_max_position_embeddings"] == 4096 + assert config.hf_text_config["rope_scaling"]["short_factor"] == (1.0,) * 48 + + +@pytest.mark.parametrize( + ("update", "message"), + [ + ({"tie_word_embeddings": False}, "tie_word_embeddings=True"), + ({"attention_bias": True}, "attention_bias=False"), + ({"hidden_act": "gelu"}, "hidden_act='silu'"), + ({"rope_scaling": {"type": "linear", "short_factor": [1.0] * 48, "long_factor": [1.0] * 48}}, "longrope"), + ({"rope_scaling": {"type": "longrope", "short_factor": [1.0] * 47, "long_factor": [1.0] * 48}}, "48 values"), + ], +) +def test_phi4mm_config_rejects_unsupported_language_semantics(update, message): + hf_config = _phi4mm_config() + hf_config.update(update) + + with pytest.raises(ValueError, match=message): + Phi4MMAdapter().config_from_hf(hf_config) + + +def test_phi4mm_registry_resolves_config(tmp_path, monkeypatch): + (tmp_path / "config.json").write_text(json.dumps(_phi4mm_config()), encoding="utf-8") + monkeypatch.setattr(registry, "_PLUGINS_LOADED", False) + monkeypatch.setattr(areno.models, "_REGISTERED_GROUPS", set()) + monkeypatch.setattr(registry, "_ADAPTERS", {}) + + config = registry.config_from_hf(tmp_path) + + assert config.model_type == "phi4mm" + assert isinstance(registry.adapter_from_hf(tmp_path), Phi4MMAdapter) + + +def test_phi4mm_tp_validation_rejects_non_divisible_kv_heads(): + config = Phi4MMAdapter().config_from_hf(_phi4mm_config()) + + config.validate_tp(1) + config.validate_tp(2) + config.validate_tp(4) + config.validate_tp(8) + with pytest.raises(ValueError, match="num_key_value_heads must be divisible"): + config.validate_tp(3) + with pytest.raises(ValueError, match="num_key_value_heads must be divisible"): + config.validate_tp(6) + + +def test_phi4mm_model_construction_has_expected_text_layers(): + config = _tiny_model_config() + model = Phi4MMAdapter().build(config) + + assert isinstance(model, Phi4MMForCausalLM) + assert len(model.model.layers) == 2 + assert model.model.embed_tokens.weight.shape == (32, 64) + assert model.lm_head.weight.shape == (32, 64) + assert model.model.norm.eps == 1e-5 + for layer in model.model.layers: + assert layer.input_layernorm.eps == 1e-5 + assert layer.post_attention_layernorm.eps == 1e-5 + assert layer.self_attn.qkv_proj.out_features == (64, 32, 32) + assert layer.self_attn.qkv_proj.local_out_features == [64, 32, 32] + assert layer.self_attn.o_proj.weight.shape == (64, 64) + assert layer.mlp.gate_up_proj.out_features == (128, 128) + assert layer.mlp.gate_up_proj.weight.shape == (256, 64) + assert layer.mlp.down_proj.weight.shape == (64, 128) + + +def test_phi4mm_projection_biases_and_qk_norm_are_disabled(): + model = Phi4MMAdapter().build(_tiny_model_config()) + + assert not hasattr(model.lm_head, "bias") + for layer in model.model.layers: + assert layer.self_attn.qkv_proj.bias is None + assert layer.self_attn.o_proj.bias is None + assert layer.self_attn.q_norm is None + assert layer.self_attn.k_norm is None + assert layer.mlp.gate_up_proj.bias is None + assert layer.mlp.down_proj.bias is None + + +def test_phi4mm_embedding_and_lm_head_share_one_optimizer_parameter(): + model = Phi4MMAdapter().build(_tiny_model_config()) + + assert model.lm_head.weight is model.model.embed_tokens.weight + parameter_ids = [id(parameter) for parameter in model.parameters()] + assert len(parameter_ids) == len(set(parameter_ids)) + + optimizer = build_optimizer( + model.parameters(), + OptimizerConfig(), + type("Context", (), {"dp_rank": 0, "dp_size": 1, "dp_group": None})(), + ) + optimizer_parameter_ids = [id(parameter) for parameter in optimizer.model_params] + assert len(optimizer_parameter_ids) == len(set(optimizer_parameter_ids)) + assert optimizer_parameter_ids.count(id(model.model.embed_tokens.weight)) == 1 + + +def test_phi4mm_text_forward_shapes_and_causal_prefix(cpu_reference_kernels): + del cpu_reference_kernels + torch.manual_seed(0) + model = Phi4MMAdapter().build(_tiny_model_config()).eval() + input_ids = torch.tensor([[1, 2, 3], [1, 2, 4]]) + + output = model(input_ids) + + assert output.hidden_states is not None + assert output.logits_shard is not None + assert output.hidden_states.shape == (2, 3, 64) + assert output.logits_shard.shape == (2, 3, 32) + assert torch.isfinite(output.hidden_states).all() + assert torch.isfinite(output.logits_shard).all() + torch.testing.assert_close(output.logits_shard[0, :2], output.logits_shard[1, :2]) + + +def test_phi4mm_lm_head_runs_inside_sequence_parallel_region(cpu_reference_kernels, monkeypatch): + del cpu_reference_kernels + model = Phi4MMAdapter().build(_tiny_model_config()).eval() + original_forward = model.lm_head.forward + sequence_parallel_states = [] + + def record_sequence_parallel_state(hidden_states): + sequence_parallel_states.append(is_sequence_parallel_active()) + return original_forward(hidden_states) + + monkeypatch.setattr(model.lm_head, "forward", record_sequence_parallel_state) + model(torch.tensor([[1, 2, 3]]), train_meta=TrainMeta(sequence_parallel=True)) + + assert sequence_parallel_states == [True] + + +def test_phi4mm_kv_cache_lifecycle(): + model = Phi4MMAdapter().build(_tiny_model_config()) + caches = model.allocate_kv_caches(num_blocks=3, block_size=4, device=torch.device("cpu")) + + assert len(caches) == len(model.layers) + assert caches[0][0].shape == (3, 4, 4, 8) + assert caches[0][0].dtype == model.config.dtype + + model.set_kv_caches(caches) + assert model.layers[0].self_attn.k_cache is caches[0][0] + assert model.onload_kv_caches(torch.device("cpu")) + + model.offload_kv_caches() + assert model.layers[0].self_attn.infer_backend is None + model.clear_kv_caches() + assert model.layers[0].self_attn.k_cache.numel() == 0 + assert not model.onload_kv_caches(torch.device("cpu")) + + with pytest.raises(ValueError, match="expected 2 layer caches"): + model.set_kv_caches(caches[:1]) + + +def _official_longrope_reference( + x: torch.Tensor, + position_ids: torch.Tensor, + dim: int, + base: float, + factors: tuple[float, ...], + max_position_embeddings: int, + original_max_position_embeddings: int, +) -> tuple[torch.Tensor, torch.Tensor]: + ext_factors = torch.tensor(factors, dtype=torch.float32, device=x.device) + inv_freq_shape = torch.arange(0, dim, 2, dtype=torch.int64, device=x.device).float() / dim + inv_freq = 1.0 / (ext_factors * base**inv_freq_shape) + inv_freq_expanded = inv_freq[None, :, None].float().expand(position_ids.shape[0], -1, 1) + position_ids_expanded = position_ids[:, None, :].float() + freqs = (inv_freq_expanded @ position_ids_expanded).transpose(1, 2) + embedding = torch.cat((freqs, freqs), dim=-1) + scale = max_position_embeddings / original_max_position_embeddings + scaling_factor = ( + 1.0 if scale <= 1.0 else math.sqrt(1.0 + math.log(scale) / math.log(original_max_position_embeddings)) + ) + return (embedding.cos() * scaling_factor).to(x.dtype), (embedding.sin() * scaling_factor).to(x.dtype) + + +def test_phi4mm_official_config_builds_partial_longrope_without_position_caches(): + config = Phi4MMAdapter().config_from_hf(_phi4mm_config()) + + rope = Phi4MMLongRoPEScaledRotaryEmbedding(config) + + assert rope.dim == 96 + assert config.head_dim - rope.dim == 32 + assert rope.short_inv_freq.shape == (48,) + assert rope.long_inv_freq.shape == (48,) + assert all("cached" not in name for name, _ in rope.named_buffers()) + + +def test_phi4mm_longrope_keeps_inverse_frequencies_in_fp32_when_model_is_cast(): + rope = Phi4MMLongRoPEScaledRotaryEmbedding(_tiny_model_config()).to(dtype=torch.bfloat16) + + assert rope.short_inv_freq.dtype == torch.float32 + assert rope.long_inv_freq.dtype == torch.float32 + + +@pytest.mark.parametrize( + ("sequence_length", "positions", "factor_key"), + [ + (32, [0, 1, 7, 31], "short_factor"), + (64, [0, 1, 31, 32, 63], "long_factor"), + ], +) +@pytest.mark.parametrize("dtype", [torch.float32, torch.bfloat16]) +def test_phi4mm_longrope_cos_sin_matches_official_reference(sequence_length, positions, factor_key, dtype): + config = _tiny_model_config() + rope = Phi4MMLongRoPEScaledRotaryEmbedding(config) + x = torch.zeros(1, len(positions), 1, config.head_dim, dtype=dtype) + position_ids = torch.tensor([positions]) + + actual_cos, actual_sin = rope.cos_sin(x, position_ids, sequence_length) + expected_cos, expected_sin = _official_longrope_reference( + x, + position_ids, + rope.dim, + config.rope_theta, + config.hf_text_config["rope_scaling"][factor_key], + config.max_position_embeddings, + config.hf_text_config["original_max_position_embeddings"], + ) + + torch.testing.assert_close(actual_cos, expected_cos, rtol=0, atol=0) + torch.testing.assert_close(actual_sin, expected_sin, rtol=0, atol=0) + + +def test_phi4mm_longrope_preserves_non_rotary_head_dimensions_and_applies_scale(): + config = _tiny_model_config() + rope = Phi4MMLongRoPEScaledRotaryEmbedding(config) + q = torch.randn(1, 3, 2, config.head_dim) + k = torch.randn(1, 3, 1, config.head_dim) + position_ids = torch.tensor([[0, 1, 2]]) + + rotated_q, rotated_k = rope(q, k, position_ids, sequence_length=3) + cos, sin = rope.cos_sin(q, torch.tensor([[0]]), sequence_length=3) + + torch.testing.assert_close(rotated_q[..., rope.dim :], q[..., rope.dim :]) + torch.testing.assert_close(rotated_k[..., rope.dim :], k[..., rope.dim :]) + assert cos[0, 0, 0].item() == pytest.approx(rope.scaling_factor) + assert sin[0, 0, 0].item() == 0.0 + + +def test_phi4mm_longrope_full_long_prefill_selects_long_factors(): + attention = Phi4MMAdapter().build(_tiny_model_config()).model.layers[0].self_attn + positions = torch.arange(40).unsqueeze(0) + infer_meta = InferMeta(mode="prefill", cu_seqlens=torch.tensor([0, 40], dtype=torch.int32), max_seqlen=40) + q = torch.randn(1, 40, attention.local_heads, attention.head_dim) + k = torch.randn(1, 40, attention.local_kv_heads, attention.head_dim) + + sequence_length = _phi4mm_longrope_sequence_length(positions, None, infer_meta, 32) + actual_q, actual_k = attention.apply_rotary(q, k, positions, None, infer_meta) + expected_q, expected_k = attention.rope(q, k, positions, sequence_length=40) + + assert sequence_length == 40 + torch.testing.assert_close(actual_q, expected_q) + torch.testing.assert_close(actual_k, expected_k) + + +def test_phi4mm_longrope_rejects_chunked_prefill_crossing_boundary(): + positions = torch.arange(28, 40).unsqueeze(0) + infer_meta = InferMeta(mode="prefill", cu_seqlens=torch.tensor([0, 12], dtype=torch.int32), max_seqlen=12) + + with pytest.raises(ValueError, match="chunked prefill cannot cross"): + _phi4mm_longrope_sequence_length(positions, None, infer_meta, 32) + + +def test_phi4mm_longrope_rejects_cached_decode_crossing_boundary(): + below_boundary = InferMeta(mode="decode", cache_seqlens=torch.tensor([31], dtype=torch.int32)) + crossing_boundary = InferMeta(mode="decode", cache_seqlens=torch.tensor([32], dtype=torch.int32)) + + assert _phi4mm_longrope_sequence_length(torch.tensor([[31]]), None, below_boundary, 32) == 32 + with pytest.raises(ValueError, match="cached decode cannot cross"): + _phi4mm_longrope_sequence_length(torch.tensor([[32]]), None, crossing_boundary, 32) + + +def test_phi4mm_longrope_decode_graph_replay_enforces_same_boundary(): + _validate_decode_cache_length(torch.tensor([31, 99], dtype=torch.int32), actual=1, limit=32) + with pytest.raises(ValueError, match="rotary-factor boundary"): + _validate_decode_cache_length(torch.tensor([32], dtype=torch.int32), actual=1, limit=32) + + +@pytest.mark.parametrize("tp_size", [1, 2, 4]) +def test_phi4mm_tp_construction_uses_compatible_local_shards(tp_size): + old_context = get_tp_context() + try: + set_tp_context(TPContext(rank=0, world_size=tp_size, device=torch.device("cpu"), group=None)) + config = _tiny_model_config() + config.validate_tp(tp_size) + model = Phi4MMAdapter().build(config) + finally: + set_tp_context(old_context) + + attention = model.model.layers[0].self_attn + assert attention.local_heads == 8 // tp_size + assert attention.local_kv_heads == 4 // tp_size + assert attention.qkv_proj.local_out_features == [64 // tp_size, 32 // tp_size, 32 // tp_size] + assert attention.o_proj.weight.shape == (64, 64 // tp_size) + assert model.model.layers[0].mlp.gate_up_proj.weight.shape == (256 // tp_size, 64) + assert model.model.layers[0].mlp.down_proj.weight.shape == (64, 128 // tp_size) + assert model.model.embed_tokens.weight.shape == (32 // tp_size, 64) + assert model.lm_head.weight.shape == (32 // tp_size, 64) + assert model.lm_head.weight is model.model.embed_tokens.weight diff --git a/tests/test_phi4mm_checkpoint_cpu.py b/tests/test_phi4mm_checkpoint_cpu.py new file mode 100644 index 00000000..456547af --- /dev/null +++ b/tests/test_phi4mm_checkpoint_cpu.py @@ -0,0 +1,275 @@ +from __future__ import annotations + +import json + +import pytest +import torch +from safetensors import safe_open +from safetensors.torch import save_file + +from areno.engine.checkpoints.common import load_packed_section_column_spec, save_packed_section_column_spec +from areno.engine.checkpoints.io import PolicyTensorStore, SafetensorsIndex +from areno.engine.config import ModelConfig +from areno.engine.parallel.context import TPContext, get_tp_context, set_tp_context +from areno.models.phi4mm import Phi4MMAdapter +from areno.models.phi4mm.checkpoint import QKV_SPEC, audit_phi4mm_checkpoint + + +@pytest.fixture(autouse=True) +def _isolate_tp_context(): + previous_context = get_tp_context() + set_tp_context(TPContext(rank=0, world_size=1, device=torch.device("cpu"), group=None)) + try: + yield + finally: + set_tp_context(previous_context) + + +def _tiny_config() -> ModelConfig: + return ModelConfig( + model_type="phi4mm", + vocab_size=32, + hidden_size=64, + intermediate_size=128, + num_hidden_layers=2, + num_attention_heads=8, + num_key_value_heads=4, + head_dim=8, + rms_norm_eps=1e-5, + rope_theta=10_000.0, + max_position_embeddings=64, + tie_word_embeddings=True, + qkv_bias=False, + qk_norm=False, + dtype=torch.float32, + hidden_act="silu", + partial_rotary_factor=0.75, + sequence_parallel=False, + attn_backend="native", + hf_text_config={ + "original_max_position_embeddings": 32, + "rope_scaling": { + "type": "longrope", + "short_factor": (1.0, 1.0, 1.0), + "long_factor": (1.0, 2.0, 3.0), + }, + }, + ) + + +def _row_values(rows: int, columns: int, base: float) -> torch.Tensor: + return (base + torch.arange(rows, dtype=torch.float32)).unsqueeze(1).expand(rows, columns).clone() + + +def _column_values(rows: int, columns: int, base: float) -> torch.Tensor: + return (base + torch.arange(columns, dtype=torch.float32)).unsqueeze(0).expand(rows, columns).clone() + + +def _synthetic_weights(*, skipped: bool = False) -> dict[str, torch.Tensor]: + config = _tiny_config() + tensors = { + "model.embed_tokens.weight": torch.arange(config.vocab_size * config.hidden_size, dtype=torch.float32).view( + config.vocab_size, config.hidden_size + ), + "model.norm.weight": torch.arange(config.hidden_size, dtype=torch.float32) + 10, + } + for layer in range(config.num_hidden_layers): + prefix = f"model.layers.{layer}" + offset = layer * 10_000 + q = _row_values(64, 64, 1_000 + offset) + k = _row_values(32, 64, 2_000 + offset) + v = _row_values(32, 64, 3_000 + offset) + gate = _row_values(128, 64, 4_000 + offset) + up = _row_values(128, 64, 5_000 + offset) + tensors.update( + { + f"{prefix}.input_layernorm.weight": torch.arange(64, dtype=torch.float32) + 20 + offset, + f"{prefix}.post_attention_layernorm.weight": torch.arange(64, dtype=torch.float32) + 30 + offset, + f"{prefix}.self_attn.qkv_proj.base_layer.weight": torch.cat((q, k, v)), + f"{prefix}.self_attn.o_proj.base_layer.weight": _column_values(64, 64, 6_000 + offset), + f"{prefix}.mlp.gate_up_proj.base_layer.weight": torch.cat((gate, up)), + f"{prefix}.mlp.down_proj.base_layer.weight": _column_values(64, 128, 7_000 + offset), + } + ) + if skipped: + tensors.update( + { + "model.layers.0.self_attn.qkv_proj.lora_A.vision.weight": torch.ones(1), + "model.layers.0.self_attn.qkv_proj.lora_B.speech.weight": torch.ones(1), + "model.embed_tokens_extend.image_embed.img_projection.weight": torch.ones(1), + "model.embed_tokens_extend.audio_embed.audio_projection.weight": torch.ones(1), + } + ) + return tensors + + +def _write_checkpoint(path, tensors: dict[str, torch.Tensor]) -> None: + path.mkdir() + save_file(tensors, path / "model.safetensors") + (path / "config.json").write_text(json.dumps({"model_type": "phi4mm"}), encoding="utf-8") + + +@pytest.mark.parametrize("tp_size", [1, 2, 4]) +def test_phi4mm_checkpoint_loads_each_packed_section_independently(tmp_path, monkeypatch, tp_size): + monkeypatch.setenv("ARENO_CKPT_PROGRESS", "0") + tensors = _synthetic_weights() + checkpoint = tmp_path / "source" + _write_checkpoint(checkpoint, tensors) + old_context = get_tp_context() + try: + for rank in range(tp_size): + set_tp_context(TPContext(rank=rank, world_size=tp_size, device=torch.device("cpu"), group=None)) + model = Phi4MMAdapter().build(_tiny_config()) + Phi4MMAdapter().load_weights(model, checkpoint) + + layer = model.model.layers[0] + q, k, v = tensors["model.layers.0.self_attn.qkv_proj.base_layer.weight"].split((64, 32, 32)) + gate, up = tensors["model.layers.0.mlp.gate_up_proj.base_layer.weight"].split((128, 128)) + expected_qkv = torch.cat((q.chunk(tp_size)[rank], k.chunk(tp_size)[rank], v.chunk(tp_size)[rank])) + expected_gate_up = torch.cat((gate.chunk(tp_size)[rank], up.chunk(tp_size)[rank])) + + torch.testing.assert_close(layer.self_attn.qkv_proj.weight, expected_qkv) + torch.testing.assert_close(layer.mlp.gate_up_proj.weight, expected_gate_up) + torch.testing.assert_close( + layer.self_attn.o_proj.weight, + tensors["model.layers.0.self_attn.o_proj.base_layer.weight"].chunk(tp_size, dim=1)[rank], + ) + torch.testing.assert_close( + layer.mlp.down_proj.weight, + tensors["model.layers.0.mlp.down_proj.base_layer.weight"].chunk(tp_size, dim=1)[rank], + ) + torch.testing.assert_close( + model.model.embed_tokens.weight, + tensors["model.embed_tokens.weight"].chunk(tp_size)[rank], + ) + torch.testing.assert_close(layer.input_layernorm.weight, tensors["model.layers.0.input_layernorm.weight"]) + torch.testing.assert_close( + layer.post_attention_layernorm.weight, + tensors["model.layers.0.post_attention_layernorm.weight"], + ) + torch.testing.assert_close(model.model.norm.weight, tensors["model.norm.weight"]) + assert model.lm_head.weight is model.model.embed_tokens.weight + finally: + set_tp_context(old_context) + + +def test_phi4mm_checkpoint_audit_accepts_only_documented_skips(tmp_path): + checkpoint = tmp_path / "source" + _write_checkpoint(checkpoint, _synthetic_weights(skipped=True)) + + audit = audit_phi4mm_checkpoint(checkpoint, num_hidden_layers=2) + + assert audit.total == 18 + assert audit.consumed == 14 + assert audit.vision_lora_skipped == 1 + assert audit.speech_lora_skipped == 1 + assert audit.vision_skipped == 1 + assert audit.audio_skipped == 1 + assert audit.unknown == 0 + + +def test_phi4mm_checkpoint_audit_rejects_unknown_base_key(tmp_path): + tensors = _synthetic_weights() + tensors["model.layers.0.self_attn.foo.weight"] = torch.ones(1) + checkpoint = tmp_path / "source" + _write_checkpoint(checkpoint, tensors) + + with pytest.raises(ValueError, match="unknown tensors.*self_attn.foo.weight"): + audit_phi4mm_checkpoint(checkpoint, num_hidden_layers=2) + + +def test_phi4mm_checkpoint_audit_rejects_missing_required_key(tmp_path): + tensors = _synthetic_weights() + del tensors["model.layers.0.self_attn.qkv_proj.base_layer.weight"] + checkpoint = tmp_path / "source" + _write_checkpoint(checkpoint, tensors) + + with pytest.raises(ValueError, match="missing 1 required.*qkv_proj.base_layer.weight"): + audit_phi4mm_checkpoint(checkpoint, num_hidden_layers=2) + + +def test_phi4mm_checkpoint_rejects_wrong_packed_shape(tmp_path, monkeypatch): + monkeypatch.setenv("ARENO_CKPT_PROGRESS", "0") + tensors = _synthetic_weights() + tensors["model.layers.0.self_attn.qkv_proj.base_layer.weight"] = torch.zeros(127, 64) + checkpoint = tmp_path / "source" + _write_checkpoint(checkpoint, tensors) + + with pytest.raises(ValueError, match=r"shape \(127, 64\), expected \(128, 64\)"): + Phi4MMAdapter().load_weights(Phi4MMAdapter().build(_tiny_config()), checkpoint) + + +def test_packed_section_loader_rejects_non_divisible_sections(tmp_path): + checkpoint = tmp_path / "source" + _write_checkpoint(checkpoint, _synthetic_weights()) + model = Phi4MMAdapter().build(_tiny_config()) + index = SafetensorsIndex(checkpoint, progress=False) + try: + with pytest.raises(ValueError, match="cannot shard size 64 across 3 ranks"): + load_packed_section_column_spec(model.model.layers[0], index, "model.layers.0", QKV_SPEC, 0, 3) + finally: + index.close() + + +@pytest.mark.parametrize("tp_size", [2, 4]) +def test_packed_section_save_layout_is_inverse_of_section_sharding(tp_size): + tensors = _synthetic_weights() + full = tensors["model.layers.0.self_attn.qkv_proj.base_layer.weight"] + q, k, v = full.split((64, 32, 32)) + old_context = get_tp_context() + contributions = [] + try: + for rank in range(tp_size): + set_tp_context(TPContext(rank=rank, world_size=tp_size, device=torch.device("cpu"), group=None)) + layer = Phi4MMAdapter().build(_tiny_config()).model.layers[0] + layer.self_attn.qkv_proj.weight.data.copy_( + torch.cat((q.chunk(tp_size)[rank], k.chunk(tp_size)[rank], v.chunk(tp_size)[rank])) + ) + store = PolicyTensorStore() + save_packed_section_column_spec(store, layer, "model.layers.0", QKV_SPEC) + layout = store["model.layers.0.self_attn.qkv_proj.base_layer.weight"].policy_layout() + contribution = torch.empty(layout.numel, dtype=layout.dtype) + layout.read_chunk(0, contribution) + contributions.append(contribution) + finally: + set_tp_context(old_context) + + reconstructed = torch.stack(contributions).sum(dim=0).reshape_as(full) + torch.testing.assert_close(reconstructed, full, rtol=0, atol=0) + + +def test_phi4mm_text_only_checkpoint_load_save_reload_closes(tmp_path, monkeypatch): + monkeypatch.setenv("ARENO_CKPT_PROGRESS", "0") + source = tmp_path / "source" + output = tmp_path / "output" + tensors = _synthetic_weights(skipped=True) + _write_checkpoint(source, tensors) + first = Phi4MMAdapter().build(_tiny_config()) + Phi4MMAdapter().load_weights(first, source) + + saved_path = Phi4MMAdapter().save_weights(first, output, source) + second = Phi4MMAdapter().build(_tiny_config()) + Phi4MMAdapter().load_weights(second, output) + + assert saved_path == str(output) + assert (output / "config.json").exists() + assert second.lm_head.weight is second.model.embed_tokens.weight + for (first_name, first_parameter), (second_name, second_parameter) in zip( + first.named_parameters(), second.named_parameters(), strict=True + ): + assert first_name == second_name + torch.testing.assert_close(first_parameter, second_parameter, rtol=0, atol=0) + audit = audit_phi4mm_checkpoint(output, num_hidden_layers=2) + assert audit.total == audit.consumed == 14 + with open(output / "model.safetensors.index.json", encoding="utf-8") as handle: + saved_keys = set(json.load(handle)["weight_map"]) + assert not any("embed_tokens_extend" in key or ".lora_" in key for key in saved_keys) + with safe_open(output / "model-rank00000-00002-layer-00000.safetensors", framework="pt") as handle: + torch.testing.assert_close( + handle.get_tensor("model.layers.0.self_attn.qkv_proj.base_layer.weight"), + tensors["model.layers.0.self_attn.qkv_proj.base_layer.weight"], + ) + torch.testing.assert_close( + handle.get_tensor("model.layers.0.mlp.gate_up_proj.base_layer.weight"), + tensors["model.layers.0.mlp.gate_up_proj.base_layer.weight"], + ) From 264b4afe8c8d44762a023fd370477a2dbff3a42a Mon Sep 17 00:00:00 2001 From: zitai-wang <2531131993@qq.com> Date: Wed, 26 Aug 2026 19:09:33 +0800 Subject: [PATCH 4/5] test(models): skip Phi-4 runtime tests without Triton --- tests/test_phi4mm_adapter_cpu.py | 2 ++ tests/test_phi4mm_checkpoint_cpu.py | 2 ++ 2 files changed, 4 insertions(+) diff --git a/tests/test_phi4mm_adapter_cpu.py b/tests/test_phi4mm_adapter_cpu.py index f4c6fcc5..e18caaed 100644 --- a/tests/test_phi4mm_adapter_cpu.py +++ b/tests/test_phi4mm_adapter_cpu.py @@ -7,6 +7,8 @@ import torch import torch.nn.functional as F +pytest.importorskip("triton") + import areno.models from areno.engine.config import ModelConfig, OptimizerConfig from areno.engine.layers import mlp, norm, vocab diff --git a/tests/test_phi4mm_checkpoint_cpu.py b/tests/test_phi4mm_checkpoint_cpu.py index 456547af..7d17fd04 100644 --- a/tests/test_phi4mm_checkpoint_cpu.py +++ b/tests/test_phi4mm_checkpoint_cpu.py @@ -7,6 +7,8 @@ from safetensors import safe_open from safetensors.torch import save_file +pytest.importorskip("triton") + from areno.engine.checkpoints.common import load_packed_section_column_spec, save_packed_section_column_spec from areno.engine.checkpoints.io import PolicyTensorStore, SafetensorsIndex from areno.engine.config import ModelConfig From 6cd583b778e5e35c0005af61f4432c8144e5adba Mon Sep 17 00:00:00 2001 From: zitai-wang <2531131993@qq.com> Date: Thu, 27 Aug 2026 10:47:57 +0800 Subject: [PATCH 5/5] fix(runtime): make CPU model imports Triton-optional --- areno/accel/ops.py | 44 +---------------------------- areno/accel/utils.py | 41 +++++++++++++++++++++++++++ areno/engine/layers/mlp.py | 3 +- areno/engine/layers/norm.py | 6 ++-- areno/models/phi4mm/model.py | 2 +- tests/test_phi4mm_adapter_cpu.py | 16 +++++++++-- tests/test_phi4mm_checkpoint_cpu.py | 2 -- 7 files changed, 63 insertions(+), 51 deletions(-) create mode 100644 areno/accel/utils.py diff --git a/areno/accel/ops.py b/areno/accel/ops.py index 6c85d7f1..76b278ea 100644 --- a/areno/accel/ops.py +++ b/areno/accel/ops.py @@ -12,11 +12,8 @@ from __future__ import annotations -import logging from typing import Any -import torch - from areno.accel.activations import areno_gelu_tanh_and_mul, areno_silu_and_mul from areno.accel.attention import ( areno_causal_attention, @@ -28,46 +25,7 @@ from areno.accel.kernels.fused_moe import is_available as fused_moe_is_available from areno.accel.kernels.group_rmsnorm import rms_norm_gate_fwd from areno.accel.kernels.seg_la import SegLaMeta, seg_la_fwd - -logger = logging.getLogger(__name__) -# Process-wide set of message keys already emitted by log_once/warn_once. -_LOGGED: set[str] = set() - - -def log_once(key: str, message: str, *, level: int = logging.DEBUG) -> None: - """Log ``message`` at most once per process for the given ``key``.""" - - if key in _LOGGED: - return - logger.log(level, message) - _LOGGED.add(key) - - -def warn_once(key: str, message: str) -> None: - """Emit a warning at most once per process for the given ``key``.""" - - log_once(key, message, level=logging.WARNING) - - -@torch._dynamo.disable -def is_cuda_graph_capturing(tensor: torch.Tensor) -> bool: - """True if the tensor lives on CUDA and we are inside a graph capture.""" - - return tensor.is_cuda and torch.cuda.is_current_stream_capturing() - - -@torch._dynamo.disable -def can_use_cuda_kernel(tensor: torch.Tensor, name: str, *, allow_sm121: bool = False) -> bool: - """Decide whether to take the fused kernel path for ``tensor``. - - Returns False only on non-CUDA tensors. ``name`` and ``allow_sm121`` are - kept for compatibility with existing call sites. - """ - - if not tensor.is_cuda: - return False - return True - +from areno.accel.utils import can_use_cuda_kernel, is_cuda_graph_capturing, log_once, warn_once __all__ = [ "Any", diff --git a/areno/accel/utils.py b/areno/accel/utils.py new file mode 100644 index 00000000..a04a5696 --- /dev/null +++ b/areno/accel/utils.py @@ -0,0 +1,41 @@ +"""Lightweight acceleration helpers that do not import optional kernels.""" + +from __future__ import annotations + +import logging + +import torch + +logger = logging.getLogger(__name__) +_LOGGED: set[str] = set() + + +def log_once(key: str, message: str, *, level: int = logging.DEBUG) -> None: + """Log ``message`` at most once per process for the given ``key``.""" + + if key in _LOGGED: + return + logger.log(level, message) + _LOGGED.add(key) + + +def warn_once(key: str, message: str) -> None: + """Emit a warning at most once per process for the given ``key``.""" + + log_once(key, message, level=logging.WARNING) + + +@torch._dynamo.disable +def is_cuda_graph_capturing(tensor: torch.Tensor) -> bool: + """True if the tensor lives on CUDA and we are inside a graph capture.""" + + return tensor.is_cuda and torch.cuda.is_current_stream_capturing() + + +@torch._dynamo.disable +def can_use_cuda_kernel(tensor: torch.Tensor, name: str, *, allow_sm121: bool = False) -> bool: + """Return whether a fused CUDA kernel can run for ``tensor``.""" + + if not tensor.is_cuda: + return False + return True diff --git a/areno/engine/layers/mlp.py b/areno/engine/layers/mlp.py index 6b785671..65c92209 100644 --- a/areno/engine/layers/mlp.py +++ b/areno/engine/layers/mlp.py @@ -10,7 +10,8 @@ import torch from torch import nn -from areno.accel.ops import areno_silu_and_mul, log_once +from areno.accel.activations import areno_silu_and_mul +from areno.accel.utils import log_once from areno.engine.config import ModelConfig from areno.engine.layers.linear import MergedColumnParallelLinear, RowParallelLinear diff --git a/areno/engine/layers/norm.py b/areno/engine/layers/norm.py index a640a96a..e765773c 100644 --- a/areno/engine/layers/norm.py +++ b/areno/engine/layers/norm.py @@ -13,7 +13,7 @@ from torch import nn from areno.accel import areno_rmsnorm -from areno.accel.ops import can_use_cuda_kernel, log_once, rms_norm_gate_fwd +from areno.accel.utils import can_use_cuda_kernel, log_once from areno.engine.layers.linear import mark_tensor_parallel_parameter @@ -90,8 +90,10 @@ def forward(self, x: torch.Tensor, gate: torch.Tensor) -> torch.Tensor: # Reshape last dim into (groups_per_rank, group_width) for the kernel. x = x.view(*shape[:-1], self.groups_per_rank, self.group_width) gate = gate.view(*shape[:-1], self.groups_per_rank, self.group_width) - if rms_norm_gate_fwd is None or not can_use_cuda_kernel(x, "fused group RMSNorm sigmoid gate kernel"): + if not can_use_cuda_kernel(x, "fused group RMSNorm sigmoid gate kernel"): raise RuntimeError("ARENO group RMSNorm sigmoid gate requires the fused CUDA kernel") + from areno.accel.kernels.group_rmsnorm import rms_norm_gate_fwd + log_once("group_rmsnorm_sigmoid_gate", "using fused group RMSNorm sigmoid gate kernel") # Flatten the leading dims into a single batch so the kernel only # sees a 3D (B, groups, width) tensor. diff --git a/areno/models/phi4mm/model.py b/areno/models/phi4mm/model.py index 45118cf3..57f972b2 100644 --- a/areno/models/phi4mm/model.py +++ b/areno/models/phi4mm/model.py @@ -14,7 +14,7 @@ import torch from torch import nn -from areno.accel.ops import is_cuda_graph_capturing +from areno.accel.utils import is_cuda_graph_capturing from areno.engine.config import ModelConfig, _parse_dtype from areno.engine.layers.attention import CausalSelfAttention from areno.engine.layers.mlp import GatedMLP diff --git a/tests/test_phi4mm_adapter_cpu.py b/tests/test_phi4mm_adapter_cpu.py index e18caaed..c2fb7e0f 100644 --- a/tests/test_phi4mm_adapter_cpu.py +++ b/tests/test_phi4mm_adapter_cpu.py @@ -2,13 +2,13 @@ import json import math +import subprocess +import sys import pytest import torch import torch.nn.functional as F -pytest.importorskip("triton") - import areno.models from areno.engine.config import ModelConfig, OptimizerConfig from areno.engine.layers import mlp, norm, vocab @@ -97,6 +97,18 @@ def _tiny_model_config() -> ModelConfig: ) +def test_phi4mm_import_does_not_require_triton(): + script = """ +import sys +sys.modules['triton'] = None +import areno.models.phi4mm +assert 'areno.accel.kernels.fused_moe' not in sys.modules +assert 'areno.accel.kernels.group_rmsnorm' not in sys.modules +assert 'areno.accel.kernels.seg_la' not in sys.modules +""" + subprocess.run([sys.executable, "-c", script], check=True) + + @pytest.fixture def cpu_reference_kernels(monkeypatch): def embedding(input_ids, weight, vocab_start, vocab_end): diff --git a/tests/test_phi4mm_checkpoint_cpu.py b/tests/test_phi4mm_checkpoint_cpu.py index 7d17fd04..456547af 100644 --- a/tests/test_phi4mm_checkpoint_cpu.py +++ b/tests/test_phi4mm_checkpoint_cpu.py @@ -7,8 +7,6 @@ from safetensors import safe_open from safetensors.torch import save_file -pytest.importorskip("triton") - from areno.engine.checkpoints.common import load_packed_section_column_spec, save_packed_section_column_spec from areno.engine.checkpoints.io import PolicyTensorStore, SafetensorsIndex from areno.engine.config import ModelConfig