Files
AstrAI/astrai/model/components/decoder_block.py
T
ViperEkura 0c1b7664c1 refactor: split infer core into subpackages by concern
- Eliminate core/ directory into cache/, runtime/, network/ subpackages plus flat modules
- Split cache.py (647 lines) into cache/{buffer,strategy,pool}.py by layer
- Add explicit ContiguousStrategy, make AllocationStrategy a real ABC
- Move TaskCacheState to cache/strategy.py, drop string forward references
- Rename api/ to network/, server.py to app.py
- Move sample.py into runtime/ alongside executor and graph
- Simplify TaskCacheManager.__init__ to single pool param
- Expose pool.strategy and pool.req_pool as public properties
- Fix KVCache import in attention_backend.py (TYPE_CHECKING guard)
- Fix steady-state decode reading uninitialized position_ids on first step
2026-08-08 23:43:05 +08:00

75 lines
2.4 KiB
Python

from dataclasses import asdict
from typing import Optional, TypedDict
import torch.nn as nn
from torch import Tensor
from astrai.inference.cache import KVCache
from astrai.model.components.attention import AttnFactory
from astrai.model.components.mlp import FFNFactory, RouterStats
from astrai.model.components.norm import RMSNorm
class DecoderOutput(TypedDict):
hidden_states: Tensor
aux_loss: Optional[Tensor]
router_stats: Optional[RouterStats]
class DecoderBlock(nn.Module):
def __init__(self, config, layer_id: int):
super().__init__()
cfg = asdict(config)
cfg.update(
dim=config.hidden_size,
dim_ffn=config.intermediate_size,
n_layers=config.num_hidden_layers,
n_heads=config.num_attention_heads,
n_kv_heads=config.num_key_value_heads,
norm_eps=config.rms_norm_eps,
down_init_std=0.02 / (2 * config.num_hidden_layers) ** 0.5,
)
self.attention = AttnFactory.create(config.attn_type, **cfg, layer_id=layer_id)
self.input_norm = RMSNorm(config.hidden_size, config.rms_norm_eps)
self.post_attention_norm = RMSNorm(config.hidden_size, config.rms_norm_eps)
ffn_type = self._resolve_ffn_type(config, layer_id)
self.mlp = FFNFactory.create(ffn_type, **cfg)
@staticmethod
def _resolve_ffn_type(config, layer_id: int) -> str:
if config.ffn_type != "moe":
return config.ffn_type
mlp_only = config.mlp_only_layers or []
if layer_id in mlp_only:
return "mlp"
if config.decoder_sparse_step > 1:
if (layer_id + 1) % config.decoder_sparse_step != 0:
return "mlp"
return "moe"
def forward(
self,
x: Tensor,
rotary_emb: Tensor,
attention_mask: Optional[Tensor] = None,
kv_cache: Optional[KVCache] = None,
is_causal: bool = False,
) -> DecoderOutput:
attn_output = self.attention(
self.input_norm(x),
rotary_emb,
attention_mask,
kv_cache,
is_causal,
)
x = attn_output + x
normalized = self.post_attention_norm(x)
mlp_output = self.mlp(normalized)
x = mlp_output["hidden_states"] + x
return {
"hidden_states": x,
"aux_loss": mlp_output["aux_loss"],
"router_stats": mlp_output.get("router_stats"),
}