""" Implementation for GPTJ architecture. TODO: add docstring """ import dataclasses from functools import partial from typing import Any, Dict, Optional # noqa: UP035 from tvm import tirx from tvm.relax.frontend import nn from tvm.relax.frontend.nn import Tensor, op from mlc_llm import op as op_ext from mlc_llm.model.model_utils import index_last_token from mlc_llm.nn import PagedKVCache, RopeMode from mlc_llm.support import logging from mlc_llm.support import tensor_parallel as tp from mlc_llm.support.config import ConfigBase from mlc_llm.support.style import bold logger = logging.getLogger(__name__) @dataclasses.dataclass class GPTJConfig(ConfigBase): """Configuration of the GPTJ model.""" vocab_size: int n_embd: int n_layer: int n_head: int layer_norm_epsilon: int rotary_dim: int activation_function: str n_inner: int = -1 rope_scaling: Optional[Dict[str, Any]] = None # noqa: UP006 context_window_size: int = 0 prefill_chunk_size: int = 0 tensor_parallel_shards: int = 1 max_batch_size: int = 1 head_dim: int = 0 kwargs: Dict[str, Any] = dataclasses.field(default_factory=dict) # noqa: UP006 def __post_init__(self): if self.context_window_size == 0: for name in ["max_position_embeddings", "n_positions"]: if name in self.kwargs: self.context_window_size = self.kwargs.pop(name) logger.info( "%s not found in config.json. Falling back to %s (%d)", bold("context_window_size"), bold(name), self.context_window_size, ) break else: raise ValueError( "Unable to determine the maximum sequence length, because none of " "`context_window_size`, `max_position_embeddings` or `max_sequence_length` is " "provided in `config.json`." ) if self.head_dim == 0: self.head_dim = self.n_embd // self.n_head assert self.head_dim * self.n_head == self.n_embd if self.prefill_chunk_size == 0: logger.info( "%s defaults to %d", bold("prefill_chunk_size"), min(self.context_window_size, 8192), ) self.prefill_chunk_size = min(self.context_window_size, 8192) elif self.prefill_chunk_size > self.context_window_size: logger.info( "Overriding %s from %d to %d", bold("prefill_chunk_size"), self.prefill_chunk_size, min(self.context_window_size, 8192), ) self.prefill_chunk_size = min(self.context_window_size, 8192) class GPTJAttention(nn.Module): def __init__(self, config: GPTJConfig): self.embed_dim = config.n_embd self.num_heads = config.n_head // config.tensor_parallel_shards self.head_dim = config.head_dim self.max_position_embeddings = config.context_window_size self.rope_theta = 10000 self.rotary_dim = config.rotary_dim self.c_attn = nn.Linear( in_features=self.embed_dim, out_features=3 * self.embed_dim, bias=False, ) self.out_proj = nn.Linear(self.embed_dim, self.embed_dim, bias=False) def forward( self, hidden_states: Tensor, paged_kv_cache: PagedKVCache, layer_id: int, ): d, h = self.head_dim, self.num_heads b, s, _ = hidden_states.shape qkv = self.c_attn(hidden_states) qkv = op.reshape(qkv, (b, s, 3 * h, d)) output = op.reshape( paged_kv_cache.attention_with_fused_qkv( layer_id, qkv, self.num_heads, sm_scale=self.head_dim**-0.5 ), (b, s, h * d), ) return self.out_proj(output) ACT2FN = { "gelu": partial(nn.gelu, approximate=False), "relu": nn.relu, "silu": nn.silu, "swish": nn.silu, "gelu_new": partial(nn.gelu, approximate=True), } class GPTJMLP(nn.Module): def __init__(self, config: GPTJConfig): # in MLP: intermediate_size= 4 * embed_dim embed_dim = config.n_embd inner_dim = 4 * config.n_embd if config.n_inner is None else config.n_inner self.fc_in = nn.Linear(embed_dim, inner_dim, bias=True) self.fc_out = nn.Linear(inner_dim, embed_dim, bias=True) self.act_fn = ACT2FN[config.activation_function] def forward(self, hidden_states: Tensor): hidden_states = self.fc_in(hidden_states) hidden_states = self.act_fn(hidden_states) hidden_states = self.fc_out(hidden_states) return hidden_states class GPTJBlock(nn.Module): def __init__(self, config: GPTJConfig): self.ln_1 = nn.LayerNorm(config.n_embd, eps=config.layer_norm_epsilon) self.attn = GPTJAttention(config) self.mlp = GPTJMLP(config) def _set_tp(): def _set(layer, hint): layer.attrs["shard_strategy"] = hint hd = config.head_dim q = self.attn.num_heads * hd k = self.attn.num_heads * hd v = self.attn.num_heads * hd _set( self.attn.c_attn.weight, tp.ShardSingleDim("_shard_qkv_weight", dim=0, segs=[q, k, v]), ) _set(self.attn.out_proj.weight, tp.ShardSingleDim("_shard_o", dim=1)) _set( self.mlp.fc_in.weight, tp.ShardSingleDim("_shard_c_fc_weight", dim=0), ) _set(self.mlp.fc_out.weight, tp.ShardSingleDim("_shard_mlp_c_proj", dim=1)) self.tensor_parallel_shards = config.tensor_parallel_shards _set_tp() def forward(self, hidden_states: Tensor, paged_kv_cache: PagedKVCache, layer_id: int): residual = hidden_states hidden_states = self.ln_1(hidden_states) attn_output = self.attn(hidden_states, paged_kv_cache, layer_id) feed_forward_hidden_states = self.mlp(hidden_states) hidden_states = self._apply_residual(attn_output + feed_forward_hidden_states, residual) return hidden_states def _apply_residual(self, out, residual): if self.tensor_parallel_shards > 1: return op.ccl_allreduce(out, "sum") + residual return out + residual class GPTJModel(nn.Module): def __init__(self, config: GPTJConfig): self.embed_dim = config.n_embd self.vocab_size = config.vocab_size self.wte = nn.Embedding(config.vocab_size, self.embed_dim) self.h = nn.ModuleList([GPTJBlock(config) for _ in range(config.n_layer)]) self.ln_f = nn.LayerNorm(self.embed_dim, eps=config.layer_norm_epsilon) def forward(self, inputs: Tensor, paged_kv_cache: PagedKVCache): hidden_states = inputs for layer_id, layer in enumerate(self.h): hidden_states = layer(hidden_states, paged_kv_cache, layer_id) hidden_states = self.ln_f(hidden_states) return hidden_states class GPTJForCausalLM(nn.Module): def __init__(self, config: GPTJConfig): self.transformer = GPTJModel(config) self.lm_head = nn.Linear(config.n_embd, config.vocab_size, dtype="float32") self.dtype = "float32" self.hidden_size = config.n_embd self.num_hidden_layers = config.n_layer self.intermediate_size = 4 * config.n_embd if config.n_inner is None else config.n_inner self.num_attention_heads = config.n_head self.rope_theta = 10000 self.rope_scaling = config.rope_scaling self.vocab_size = config.vocab_size self.tensor_parallel_shards = config.tensor_parallel_shards self.head_dim = config.n_embd // config.n_head self.rotary_dim = config.rotary_dim def to(self, dtype: Optional[str] = None): super().to(dtype=dtype) if dtype is not None: self.dtype = dtype def batch_forward( self, input_embeds: Tensor, paged_kv_cache: PagedKVCache, logit_positions: Optional[Tensor] = None, ): op_ext.configure() hidden_states = self.transformer(input_embeds, paged_kv_cache) if logit_positions is not None: hidden_states = op.take(hidden_states, logit_positions, axis=1) logits = self.lm_head(hidden_states) if logits.dtype != "float32": logits = logits.astype("float32") return logits def embed(self, input_ids: Tensor): if self.tensor_parallel_shards > 1: input_ids = op.ccl_broadcast_from_worker0(input_ids) return self.transformer.wte(input_ids) def prefill(self, input_embed: Tensor, paged_kv_cache: PagedKVCache): op_ext.configure() hidden_states = self.transformer(input_embed, paged_kv_cache) hidden_states = index_last_token(hidden_states) logits = self.lm_head(hidden_states) if logits.dtype != "float32": logits = logits.astype("float32") return logits, paged_kv_cache def decode(self, input_embed: Tensor, paged_kv_cache: PagedKVCache): op_ext.configure() hidden_states = self.transformer(input_embed, paged_kv_cache) logits = self.lm_head(hidden_states) if logits.dtype != "float32": logits = logits.astype("float32") return logits, paged_kv_cache def batch_prefill( self, input_embeds: Tensor, logit_positions: Tensor, paged_kv_cache: PagedKVCache, ): if self.tensor_parallel_shards > 1: logit_positions = op.ccl_broadcast_from_worker0(logit_positions) logits = self.batch_forward(input_embeds, paged_kv_cache, logit_positions) return logits, paged_kv_cache def batch_decode(self, input_embeds: Tensor, paged_kv_cache: PagedKVCache): logits = self.batch_forward(input_embeds, paged_kv_cache) return logits, paged_kv_cache def batch_verify(self, input_embeds: Tensor, paged_kv_cache: PagedKVCache): logits = self.batch_forward(input_embeds, paged_kv_cache) return logits, paged_kv_cache def create_paged_kv_cache( self, max_batch_size: tirx.Var, max_total_seq_len: tirx.Var, prefill_chunk_size: tirx.Var, page_size: tirx.Var, support_sliding_window: tirx.Var, ) -> PagedKVCache: return PagedKVCache.create_generic( attn_kind="mha", max_batch_size=max_batch_size, max_total_seq_len=max_total_seq_len, prefill_chunk_size=prefill_chunk_size, page_size=page_size, support_sliding_window=support_sliding_window, num_hidden_layers=self.num_hidden_layers, num_attention_heads=self.num_attention_heads // self.tensor_parallel_shards, num_key_value_heads=self.num_attention_heads // self.tensor_parallel_shards, qk_head_dim=self.head_dim, v_head_dim=self.head_dim, rope_mode=RopeMode.NORMAL, rope_scale=1, rope_theta=self.rope_theta, rotary_dim=self.rotary_dim, rope_scaling=self.rope_scaling, dtype=self.dtype, ) def get_default_spec(self): mod_spec = { "embed": { "input_ids": nn.spec.Tensor(["seq_len"], "int32"), "$": { "param_mode": "packed", "effect_mode": "none", }, }, "prefill": { "input_embed": nn.spec.Tensor([1, "seq_len", self.hidden_size], self.dtype), "paged_kv_cache": nn.spec.Object(object_type=PagedKVCache), "$": { "param_mode": "packed", "effect_mode": "none", }, }, "decode": { "input_embed": nn.spec.Tensor([1, 1, self.hidden_size], self.dtype), "paged_kv_cache": nn.spec.Object(object_type=PagedKVCache), "$": { "param_mode": "packed", "effect_mode": "none", }, }, "batch_prefill": { "input_embeds": nn.spec.Tensor([1, "seq_len", self.hidden_size], self.dtype), "logit_positions": nn.spec.Tensor(["batch_size"], "int32"), "paged_kv_cache": nn.spec.Object(object_type=PagedKVCache), "$": { "param_mode": "packed", "effect_mode": "none", }, }, "batch_decode": { "input_embeds": nn.spec.Tensor(["batch_size", 1, self.hidden_size], self.dtype), "paged_kv_cache": nn.spec.Object(object_type=PagedKVCache), "$": { "param_mode": "packed", "effect_mode": "none", }, }, "batch_verify": { "input_embeds": nn.spec.Tensor([1, "seq_len", self.hidden_size], self.dtype), "paged_kv_cache": nn.spec.Object(object_type=PagedKVCache), "$": { "param_mode": "packed", "effect_mode": "none", }, }, "create_paged_kv_cache": { "max_batch_size": int, "max_total_seq_len": int, "prefill_chunk_size": int, "page_size": int, "support_sliding_window": int, "$": { "param_mode": "none", "effect_mode": "none", }, }, } return nn.spec.ModuleSpec.from_raw(mod_spec, self)