diff --git a/Engine/backend_spec.py b/Engine/backend_spec.py new file mode 100644 index 00000000..8a45f159 --- /dev/null +++ b/Engine/backend_spec.py @@ -0,0 +1,250 @@ +import torch +from MagicDec.Engine.model_spec import Transformer +from MagicDec.Engine.utils import load_model +import flashinfer + +############# utility lib functions ################ + + + +#################################################### + + +class LMBackend: + def __init__(self, dtype = torch.bfloat16, device: str = "cuda:0", dec_len: int = 1) -> None: + self.dtype = dtype + self.device = device + self.dec_len = dec_len + self.model_forward = lambda model, x, input_pos, kv_append_indptr, kv_page_indices, kv_page_indptr, kv_page_lastlen: model(x, input_pos, kv_append_indptr, kv_page_indices, kv_page_indptr, kv_page_lastlen) + self.prefill = lambda model, x, input_pos, kv_append_indptr, kv_page_indices, kv_page_indptr, kv_page_lastlen, is_last=None: model.prefill(x, input_pos, kv_append_indptr, kv_page_indices, kv_page_indptr, kv_page_lastlen, is_last) + self.cachelens = None + self.is_draft = False + + def load_model(self, checkpoints: str, use_tp: bool, rank_group=None, group = None): + self.model: Transformer = load_model(checkpoint_path=checkpoints, device=self.device, precision=self.dtype, use_tp=use_tp, rank_group=rank_group, group=group) + + @torch.inference_mode() + def setup_caches(self, max_batch_size: int = 1, max_seq_length: int = 2048): + self.max_length = max_seq_length + self.batch_size = max_batch_size + self.cachelens = torch.zeros(max_batch_size, dtype=torch.int32, device=self.device) + self.page_size = page_size = 128 + self.max_num_pages = max_num_pages = max_batch_size * (max_seq_length + page_size - 1) // page_size + self.max_num_pages_per_request = max_num_pages // max_batch_size + self.num_pages_per_request = torch.zeros(max_batch_size, dtype=torch.int32, device=self.device) + + # Init Attention Backend (Flashinfer) + self.decode_buffer = torch.empty(128 * 1024 * 1024, dtype=torch.uint8, device=self.device) + self.prefill_buffer = torch.empty(128 * 1024 * 1024, dtype=torch.uint8, device=self.device) + + self.qo_indptr = torch.arange(max_batch_size+1, dtype=torch.int32, device=self.device) + self.paged_kv_indptr = torch.arange(max_batch_size+1, dtype=torch.int32, device=self.device) + self.paged_kv_indices = torch.empty(max_num_pages, dtype=torch.int32, device=self.device) + self.paged_kv_last_page_len = torch.zeros((max_batch_size), dtype=torch.int32, device=self.device) + self.decode_wrapper = flashinfer.BatchPrefillWithPagedKVCacheWrapper(self.decode_buffer, "NHD", use_cuda_graph=True, + qo_indptr_buf=self.qo_indptr, + paged_kv_indptr_buf=self.paged_kv_indptr, + paged_kv_indices_buf=self.paged_kv_indices, + paged_kv_last_page_len_buf=self.paged_kv_last_page_len) + + self.prefill_wrapper = flashinfer.BatchPrefillWithPagedKVCacheWrapper(self.prefill_buffer, "NHD") + + torch.library.define( + "mylib::decode", + "(Tensor q, Tensor kv_cache) -> Tensor", + ) + + @torch.library.impl("mylib::decode", "cuda") + def decode(q, kv_cache): + return self.decode_wrapper.run( + q, kv_cache + ) + + @torch.library.register_fake("mylib::decode") + def decode_abstract(q, kv_cache): + return torch.empty_like(q) + + torch.library.define( + "mylib::prefill", + "(Tensor q, Tensor kv_cache) -> Tensor", + ) + + @torch.library.impl("mylib::prefill", "cuda") + def prefill(q, kv_cache): + return self.prefill_wrapper.run( + q, kv_cache + ) + + @torch.library.register_fake("mylib::prefill") + def prefill_abstract(q, kv_cache): + return torch.empty_like(q) + + with torch.device(self.device): + self.model.setup_caches(num_pages=max_num_pages, page_size=page_size) + + def compile(self): + import torch._dynamo.config + import torch._inductor.config + torch._inductor.config.coordinate_descent_tuning = True + torch._inductor.config.triton.unique_kernel_names = True + torch._inductor.config.fx_graph_cache = True + torch._functorch.config.enable_autograd_cache = True + self.model_forward = torch.compile(self.model_forward, mode="max-autotune", fullgraph=True) + + # used for both speculation, verification and autoregressive decoding + @torch.inference_mode() + def inference(self, input_ids: torch.LongTensor, benchmark = False): + dec_len = input_ids.shape[1] + self.pre_infer(dec_len=dec_len) + + logits = self.model_forward( + model=self.model, + x=input_ids, + input_pos=self.cachelens, + kv_append_indptr = self.qo_indptr*dec_len, kv_page_indices = self.paged_kv_indices, kv_page_indptr= self.paged_kv_indptr, kv_page_lastlen = self.paged_kv_last_page_len) + + self.cachelens += dec_len + if benchmark: + # If benchmarking the latency, don't update the cachelens and page table + self.cachelens -= dec_len + self.paged_kv_last_page_len -= dec_len + return logits + + def pre_infer(self, dec_len): + self.paged_kv_last_page_len += dec_len + self.decode_wrapper.plan( + qo_indptr=self.qo_indptr*dec_len, + paged_kv_indptr=self.paged_kv_indptr, + paged_kv_indices=self.paged_kv_indices, + paged_kv_last_page_len=self.paged_kv_last_page_len, + num_qo_heads=self.model.config.n_head, + num_kv_heads=self.model.config.n_local_heads, + head_dim=self.model.config.head_dim, + page_size=self.page_size, + q_data_type=self.dtype, + causal=True, + ) + + @torch.inference_mode() + def encode(self, input_ids: torch.LongTensor, benchmark = False): + self.clear_kv() + logits = None + seq_len = input_ids.shape[1] + chunk_size = 128 + num_chunks = (seq_len + chunk_size - 1) // chunk_size # Ceil division + is_last = False + for i in range(num_chunks): + start_idx = i * chunk_size + end_idx = min((i + 1) * chunk_size, seq_len) + chunk_input_ids = input_ids[:, start_idx:end_idx] + dec_len = end_idx-start_idx + is_last = (i == num_chunks-1) + self.pre_encode(dec_len=dec_len) + logits = self.prefill( + model=self.model, + x=chunk_input_ids, + input_pos=self.cachelens, + kv_append_indptr = self.qo_indptr*dec_len, + kv_page_indices = self.paged_kv_indices, + kv_page_indptr= self.paged_kv_indptr, + kv_page_lastlen = self.paged_kv_last_page_len, + is_last=(self.is_draft and is_last) + ) + self.cachelens += dec_len + + return logits + + def pre_encode(self, dec_len): + self.num_pages_per_request+=1 + qo_indptr = self.qo_indptr*dec_len + self.paged_kv_indices = torch.cat([torch.arange(i * self.max_num_pages_per_request, i * self.max_num_pages_per_request + self.num_pages_per_request[i], dtype=torch.int32, device=self.device) for i in range(self.batch_size)]) + self.paged_kv_indptr[1:] = torch.cumsum(self.num_pages_per_request, dim=0, dtype=torch.int32) + self.paged_kv_last_page_len = torch.full((self.batch_size,), dec_len, dtype=torch.int32, device=self.device) + self.prefill_wrapper.plan( + qo_indptr=qo_indptr, + paged_kv_indptr=self.paged_kv_indptr, + paged_kv_indices=self.paged_kv_indices, + paged_kv_last_page_len=self.paged_kv_last_page_len, + num_qo_heads=self.model.config.n_head, + num_kv_heads=self.model.config.n_local_heads, + head_dim=self.model.config.head_dim, + page_size=self.page_size, + q_data_type=self.dtype, + causal=True + ) + + @torch.inference_mode() + def clear_kv(self): + for b in self.model.layers: + b.attention.kv_cache.kv_cache.zero_() + self.cachelens.zero_() + self.qo_indptr = torch.arange(self.batch_size+1, dtype=torch.int32, device=self.device) + self.paged_kv_indptr = torch.arange(self.batch_size+1, dtype=torch.int32, device=self.device) + self.paged_kv_indices = torch.empty(self.max_num_pages, dtype=torch.int32, device=self.device) + self.paged_kv_last_page_len = torch.zeros((self.batch_size), dtype=torch.int32, device=self.device) + self.num_pages_per_request = torch.zeros(self.batch_size, device=self.device, dtype=torch.int32) + + +class LMBackendDraft(LMBackend): + def __init__(self, *args, **kwargs) -> None: + super().__init__(*args, **kwargs) + self.is_draft = True + + @torch.inference_mode() + def setup_caches(self, max_batch_size: int = 1, max_seq_length: int = 2048, window_size: int = 32): + self.max_length = max_seq_length + self.batch_size = max_batch_size + self.cachelens = torch.zeros(max_batch_size, dtype=torch.int32, device=self.device) + self.page_size = page_size = 128 + self.max_num_pages = max_num_pages = max_batch_size * (max_seq_length + page_size - 1) // page_size + self.max_num_pages_per_request = max_num_pages // max_batch_size + self.num_pages_per_request = torch.zeros(max_batch_size, dtype=torch.int32, device=self.device) + + # Init Attention Backend (Flashinfer) + self.decode_buffer = torch.empty(128 * 1024 * 1024, dtype=torch.uint8, device=self.device) + self.prefill_buffer = torch.empty(128 * 1024 * 1024, dtype=torch.uint8, device=self.device) + + self.qo_indptr = torch.arange(max_batch_size+1, dtype=torch.int32, device=self.device) + self.paged_kv_indptr = torch.arange(max_batch_size+1, dtype=torch.int32, device=self.device) + self.paged_kv_indices = torch.empty(max_num_pages, dtype=torch.int32, device=self.device) + self.paged_kv_last_page_len = torch.zeros((max_batch_size), dtype=torch.int32, device=self.device) + self.decode_wrapper = flashinfer.BatchPrefillWithPagedKVCacheWrapper(self.decode_buffer, "NHD", use_cuda_graph=True, + qo_indptr_buf=self.qo_indptr, + paged_kv_indptr_buf=self.paged_kv_indptr, + paged_kv_indices_buf=self.paged_kv_indices, + paged_kv_last_page_len_buf=self.paged_kv_last_page_len) + + self.prefill_wrapper = flashinfer.BatchPrefillWithPagedKVCacheWrapper(self.prefill_buffer, "NHD") + + torch.library.define( + "mylib::speculate", + "(Tensor q, Tensor kv_cache) -> Tensor", + ) + + @torch.library.impl("mylib::speculate", "cuda") + def speculate(q, kv_cache): + return self.decode_wrapper.run( + q, kv_cache + ) + + @torch.library.register_fake("mylib::speculate") + def speculate_abstract(q, kv_cache): + return torch.empty_like(q) + + torch.library.define( + "mylib::draft_prefill", + "(Tensor q, Tensor kv_cache) -> Tensor", + ) + + @torch.library.impl("mylib::draft_prefill", "cuda") + def draft_prefill(q, kv_cache): + return self.prefill_wrapper.run( + q, kv_cache + ) + + @torch.library.register_fake("mylib::draft_prefill") + def draft_prefill_abstract(q, kv_cache): + return torch.empty_like(q) + + with torch.device(self.device): + self.model.setup_caches(num_pages=max_num_pages, page_size=page_size, is_draft=True, window_size=window_size) \ No newline at end of file diff --git a/Engine/model_spec.py b/Engine/model_spec.py new file mode 100644 index 00000000..ea5e2fce --- /dev/null +++ b/Engine/model_spec.py @@ -0,0 +1,281 @@ +from dataclasses import dataclass +from typing import Optional + +from einops import rearrange +import torch +import torch.nn as nn +from torch import Tensor +from torch.nn import functional as F +import torch.distributed as dist +import math + +def find_multiple(n: int, k: int) -> int: + if n % k == 0: + return n + return n + k - (n % k) + +@dataclass +class ModelArgs: + block_size: int = 2048 + vocab_size: int = 32000 + n_layer: int = 32 + n_head: int = 32 + dim: int = 4096 + intermediate_size: int = None + n_local_heads: int = -1 + head_dim: int = 64 + rope_base: float = 10000 + norm_eps: float = 1e-5 + scaling_factor:float = 1.0 + # llama 3.1 with high_freq_factor and low_freq_factor + low_freq_factor: int = None # added new + high_freq_factor: int = None # added new + original_max_position_embeddings: int = None # added new + + def __post_init__(self): + if self.n_local_heads == -1: + self.n_local_heads = self.n_head + if self.intermediate_size is None: + hidden_dim = 4 * self.dim + n_hidden = int(2 * hidden_dim / 3) + self.intermediate_size = find_multiple(n_hidden, 256) + self.head_dim = self.dim // self.n_head + + @classmethod + def from_name(cls, name: str): + if name in transformer_configs: + return cls(**transformer_configs[name]) + # fuzzy search + config = [config for config in transformer_configs if config.lower() in str(name).lower()] + # We may have two or more configs matched (e.g. "7B" and "Mistral-7B"). Find the best config match, + # take longer name (as it have more symbols matched) + if len(config) > 1: + config.sort(key=len, reverse=True) + assert len(config[0]) != len(config[1]), name # make sure only one 'best' match + print(config) + return cls(**transformer_configs[config[0]]) + + +transformer_configs = { + "llama-2-7b": dict(block_size=4096, n_layer=32, n_head=32, dim=4096), + 'llama-2-7b-32k': dict(block_size=32768, n_layer=32, dim= 4096, vocab_size=32000, scaling_factor=8), + "llama-2-13b": dict(block_size=4096, n_layer=40, n_head=40, dim=5120), + "llama-2-70b": dict(block_size=4096, n_layer=80, n_head=64, dim=8192, n_local_heads=8, intermediate_size=28672), + "llama-3-8b": dict(block_size=8192, n_layer=32, n_head=32, n_local_heads=8, dim=4096, intermediate_size=14336, vocab_size=128256, rope_base=500000), + "llama-3-70b": dict(block_size=8192, n_layer=80, n_head=64, n_local_heads=8, dim=8192, intermediate_size=28672, vocab_size=128256, rope_base=500000), + "68m": dict(block_size=2048, n_layer=2, n_head=12, n_local_heads=12, dim=768, intermediate_size=3072, vocab_size=32000), + "tinyllama": dict(block_size =2048, n_layer=22, n_head=32, n_local_heads=4, dim=2048, intermediate_size=5632, vocab_size=32000), + "llama-3.1-8b": dict(block_size=131072, n_layer=32, n_head=32, n_local_heads=8, dim=4096, intermediate_size=14336, vocab_size=128256, rope_base=500000.0, scaling_factor=8, high_freq_factor=4, low_freq_factor=1, original_max_position_embeddings=8192), + "llama-3.1-70b": dict(block_size=131072, n_layer=80, n_head=64, n_local_heads=8, dim=8192, intermediate_size=28672, vocab_size=128256, rope_base=500000.0, scaling_factor=8, high_freq_factor=4, low_freq_factor=1, original_max_position_embeddings=8192), +} + +class KVCache(nn.Module): + def __init__(self, max_num_pages, page_size, n_heads, head_dim, dtype=torch.bfloat16): + super().__init__() + cache_shape = (max_num_pages, 2, page_size, n_heads, head_dim) + self.register_buffer('kv_cache', torch.zeros(cache_shape, dtype=dtype)) + + def update(self, k, v, kv_append_indptr, kv_page_indices, kv_page_indptr, kv_page_lastlen): + torch.ops.mylib.update_kv( + k, + v, + kv_append_indptr, + self.kv_cache, + kv_page_indices, + kv_page_indptr, + kv_page_lastlen, + ) + return self.kv_cache + + +class Transformer(nn.Module): + def __init__(self, config: ModelArgs) -> None: + super().__init__() + self.config = config + self.tok_embeddings = nn.Embedding(config.vocab_size, config.dim) + self.layers = nn.ModuleList(TransformerBlock(config) for _ in range(config.n_layer)) + self.norm = RMSNorm(config.dim, eps=config.norm_eps) + self.output = nn.Linear(config.dim, config.vocab_size, bias=False) + self.world_size = None + self.rank = None + self.process_group = None + + def setup_caches(self, num_pages, page_size, is_draft=False, budget = -1, window_size = 32, pooling="avgpool", kernel_size=5): + + head_dim = self.config.dim // self.config.n_head + # dtype = self.output.weight.dtype + dtype = self.output.weight.dtype if self.output.weight.dtype == torch.float16 else torch.bfloat16 + + for b in self.layers: + b.attention.kv_cache = KVCache(num_pages, page_size, self.config.n_local_heads, head_dim, dtype) + b.attention.attn_decode = torch.ops.mylib.decode if not is_draft else torch.ops.mylib.speculate + b.attention.attn_prefill = torch.ops.mylib.prefill if not is_draft else torch.ops.mylib.draft_prefill + b.attention.rope = torch.ops.mylib.llama31rope + # if is_draft and budget > 0: + # b.attention.is_sparse = True + # b.attention.budget = budget + # b.attention.window_size = window_size + # b.attention.pooling = pooling + # b.attention.kernel_size = kernel_size + + def forward(self, idx: Tensor, input_pos: Tensor, kv_append_indptr: Tensor, kv_page_indices: Tensor, kv_page_indptr: Tensor, kv_page_lastlen: Tensor) -> Tensor: + x = self.tok_embeddings(idx) + for i, layer in enumerate(self.layers): + x = layer(x, input_pos, kv_append_indptr, kv_page_indices, kv_page_indptr, kv_page_lastlen) + x = self.norm(x) + logits = self.output(x) + if self.process_group != None: + all_max_value = torch.zeros((x.shape[0], x.shape[1], self.world_size), dtype=logits.dtype, device=logits.device) + all_max_indices = torch.zeros((x.shape[0], x.shape[1], self.world_size), dtype=torch.long, device=logits.device) + all_max_value[:, :, self.rank], all_max_indices[:, :, self.rank] = torch.max(logits, dim=-1) + all_max_indices[:, :, self.rank] += self.rank * logits.shape[-1] + dist.all_reduce(all_max_value) + dist.all_reduce(all_max_indices) + global_select_indices = torch.argmax(all_max_value, dim=-1) + global_indices = torch.gather(all_max_indices, dim=-1, index=global_select_indices.unsqueeze(-1)) + return global_indices.squeeze(-1) + return torch.argmax(logits, dim=-1) + + def prefill(self, idx: Tensor, input_pos: Tensor, kv_append_indptr: Tensor, kv_page_indices: Tensor, kv_page_indptr: Tensor, kv_page_lastlen: Tensor, is_last = False): + x = self.tok_embeddings(idx) + for i, layer in enumerate(self.layers): + x = layer.prefill(x, input_pos, kv_append_indptr, kv_page_indices, kv_page_indptr, kv_page_lastlen, is_last) + x = self.norm(x) + logits = self.output(x) + if self.process_group != None: + all_max_value = torch.zeros((x.shape[0], x.shape[1], self.world_size), dtype=logits.dtype, device=logits.device) + all_max_indices = torch.zeros((x.shape[0], x.shape[1], self.world_size), dtype=torch.long, device=logits.device) + all_max_value[:, :, self.rank], all_max_indices[:, :, self.rank] = torch.max(logits, dim=-1) + all_max_indices[:, :, self.rank] += self.rank * logits.shape[-1] + dist.all_reduce(all_max_value) + dist.all_reduce(all_max_indices) + global_select_indices = torch.argmax(all_max_value, dim=-1) + global_indices = torch.gather(all_max_indices, dim=-1, index=global_select_indices.unsqueeze(-1)) + return global_indices.squeeze(-1) + return torch.argmax(logits, dim=-1) + + @classmethod + def from_name(cls, name: str): + return cls(ModelArgs.from_name(name)) + + +class TransformerBlock(nn.Module): + def __init__(self, config: ModelArgs) -> None: + super().__init__() + self.attention = Attention(config) + self.feed_forward = FeedForward(config) + self.ffn_norm = RMSNorm(config.dim, config.norm_eps) + self.attention_norm = RMSNorm(config.dim, config.norm_eps) + + def forward(self, x: Tensor, offsets: Tensor, kv_append_indptr: Tensor, kv_page_indices: Tensor, kv_page_indptr: Tensor, kv_page_lastlen: Tensor) -> Tensor: + h = x + self.attention(self.attention_norm(x), offsets, kv_append_indptr, kv_page_indices, kv_page_indptr, kv_page_lastlen) + out = h + self.feed_forward(self.ffn_norm(h)) + return out + + def prefill(self, x: Tensor, offsets: Tensor, kv_append_indptr: Tensor, kv_page_indices: Tensor, kv_page_indptr: Tensor, kv_page_lastlen: Tensor, is_last = False): + h = x + self.attention.prefill(self.attention_norm(x), offsets, kv_append_indptr, kv_page_indices, kv_page_indptr, kv_page_lastlen, is_last) + out = h + self.feed_forward(self.ffn_norm(h)) + return out + +class Attention(nn.Module): + def __init__(self, config: ModelArgs): + super().__init__() + assert config.dim % config.n_head == 0 + + total_head_dim = (config.n_head + 2 * config.n_local_heads) * config.head_dim + # key, query, value projections for all heads, but in a batch + self.wqkv = nn.Linear(config.dim, total_head_dim, bias=False) + self.wo = nn.Linear(config.dim, config.dim, bias=False) + self.kv_cache = None + self.process_group = None + self.attn_decode = None + self.attn_prefill = None + self.attn_draft = None + self.rope = None + self.is_sparse = False + + self.window_size = None + self.pooling = None + self.kernel_size = None + self.draft_budget = None + + self.n_head = config.n_head + self.head_dim = config.head_dim + self.n_local_heads = config.n_local_heads + self.dim = config.dim + self._register_load_state_dict_pre_hook(self.load_hook) + + def load_hook(self, state_dict, prefix, *args): + if prefix + "wq.weight" in state_dict: + wq = state_dict.pop(prefix + "wq.weight") + wk = state_dict.pop(prefix + "wk.weight") + wv = state_dict.pop(prefix + "wv.weight") + state_dict[prefix + "wqkv.weight"] = torch.cat([wq, wk, wv]) + + def forward(self, x: Tensor, offsets: Tensor, kv_append_indptr: Tensor, kv_page_indices: Tensor, kv_page_indptr: Tensor, kv_page_lastlen: Tensor) -> Tensor: + bsz, seqlen, _ = x.shape + kv_size = self.n_local_heads * self.head_dim + q, k, v = self.wqkv(x).split([self.dim, kv_size, kv_size], dim=-1) + q = q.view(bsz * seqlen, self.n_head, self.head_dim) + k = k.view(bsz * seqlen, self.n_local_heads, self.head_dim) + v = v.contiguous().view(bsz * seqlen, self.n_local_heads, self.head_dim) + q, k = self.rope(q, k, kv_append_indptr, offsets) + kv_cache = self.kv_cache.update(k, v, kv_append_indptr, kv_page_indices, kv_page_indptr, kv_page_lastlen) + y = self.attn_decode(q, kv_cache) + y = y.contiguous().view(bsz, seqlen, self.dim) + y = self.wo(y) + if self.process_group != None: + dist.all_reduce(y) + return y + + def prefill(self, x: Tensor, offsets: Tensor, kv_append_indptr: Tensor, kv_page_indices: Tensor, kv_page_indptr, kv_page_lastlen: Tensor, is_last = False): + bsz, seqlen, _ = x.shape + kv_size = self.n_local_heads * self.head_dim + q, k, v = self.wqkv(x).split([self.dim, kv_size, kv_size], dim=-1) + q = q.view(bsz * seqlen, self.n_head, self.head_dim) + k = k.view(bsz * seqlen, self.n_local_heads, self.head_dim) + v = v.contiguous().view(bsz * seqlen, self.n_local_heads, self.head_dim) + q, k = self.rope(q, k, kv_append_indptr, offsets) + kv_cache = self.kv_cache.update(k, v, kv_append_indptr, kv_page_indices, kv_page_indptr, kv_page_lastlen) + y = self.attn_prefill(q, kv_cache) + + if is_last and self.is_sparse: + self.gen_sparse_kv(q, kv_cache[:, 0], kv_cache[:, 1], bsz, seqlen, offsets[0]+seqlen, + (kv_append_indptr/seqlen*self.budget).to(torch.int32), kv_page_indptr, kv_page_indices, kv_page_lastlen) + + y = y.contiguous().view(bsz, seqlen, self.dim) + y = self.wo(y) + if self.process_group != None: + dist.all_reduce(y) + return y + + def gen_draft_kv(self, q, k, v, bsz, seqlen, context_len, kv_append_indptr): + raise NotImplementedError + +class FeedForward(nn.Module): + def __init__(self, config: ModelArgs) -> None: + super().__init__() + self.w1 = nn.Linear(config.dim, config.intermediate_size, bias=False) + self.w3 = nn.Linear(config.dim, config.intermediate_size, bias=False) + self.w2 = nn.Linear(config.intermediate_size, config.dim, bias=False) + self.process_group = None + + def forward(self, x: Tensor) -> Tensor: + y = self.w2(F.silu(self.w1(x)) * self.w3(x)) + if self.process_group != None: + dist.all_reduce(y) + return y + + +class RMSNorm(nn.Module): + def __init__(self, dim: int, eps: float = 1e-5): + super().__init__() + self.eps = eps + self.weight = nn.Parameter(torch.ones(dim)) + + def _norm(self, x): + return x * torch.rsqrt(torch.mean(x * x, dim=-1, keepdim=True) + self.eps) + + def forward(self, x: Tensor) -> Tensor: + output = self._norm(x.float()).type_as(x) + return output * self.weight \ No newline at end of file diff --git a/Engine/utils.py b/Engine/utils.py index 3adc9faa..cf1cb4b9 100644 --- a/Engine/utils.py +++ b/Engine/utils.py @@ -206,7 +206,7 @@ def setup_seed(seed): torch.backends.cudnn.deterministic = True def load_model(checkpoint_path, device, precision, use_tp, rank_group=None, group=None): - from MagicDec.Engine.model import Transformer + from MagicDec.Engine.model_spec import Transformer with torch.device('meta'): model = Transformer.from_name(checkpoint_path.parent.name) diff --git a/tests/smalldraft_benchmark.py b/tests/smalldraft_benchmark.py new file mode 100644 index 00000000..190af422 --- /dev/null +++ b/tests/smalldraft_benchmark.py @@ -0,0 +1,266 @@ +import time +import torch +import sys +sys.path.append("..") +from pathlib import Path +import torch.distributed as dist +from MagicDec.Engine.utils import setup_seed, cuda_graph_for_sampling_argmax_batch, sampling_argmax_batch +from MagicDec.Data.data_converter import convert_pg19_dataset +from transformers import AutoTokenizer +from torch.utils.data.dataloader import DataLoader +from tqdm import tqdm +import argparse +from MagicDec.Engine.backend_spec import LMBackend, LMBackendDraft + + +parser = argparse.ArgumentParser(description='Process model configuration and partitions.') +parser.add_argument('--model', type=Path, default=Path("/home/rsadhukh/opensource/specdec/gpt-fast/checkpoints/meta-llama/Meta-Llama-3.1-8B/model.pth"), help='model') +parser.add_argument('--draft', type=Path, default=Path("/home/rsadhukh/opensource/specdec/gpt-fast/checkpoints/meta-llama/Meta-Llama-3.1-8B/model.pth"), help='draft model') +parser.add_argument('--model_name', type=str, default="meta-llama/Meta-Llama-3.1-8B", help='model name') +parser.add_argument('--dataset', type=str, default="pg19", help='Dataset name.') +parser.add_argument('--draft_budget', type=int, default=-1, help='Dataset end index.') +parser.add_argument('--rank_group', nargs='+', type=int, help='Target group of ranks') +parser.add_argument('--compile', action='store_true', help='Whether to compile the model.') + +parser.add_argument('--gamma', type=int, default=7, help='start') + +parser.add_argument('--B', type=int, default=4, help='Batch size.') +parser.add_argument('--prefix_len', type=int, default=4128, help='Prefix length') +parser.add_argument('--max_len', type=int, default=4224, help='Generate length') +parser.add_argument('--window_size', type=int, default=32, help='Generate length') + +parser.add_argument('--seed', type=int, default=123, help='Random seed.') + +parser.add_argument('--printoutput', action='store_true', help='Whether to compile the model.') +parser.add_argument('--benchmark', action='store_true', help='Whether to compile the model.') + +args = parser.parse_args() +assert args.prefix_len < args.max_len +assert (args.prefix_len - args.window_size) % 128 == 0 +assert args.max_len % 128 == 0 +if args.draft_budget > 0: + assert (args.draft_budget - 1) % 128 == 0 + +# Init model parallelism +DEVICE = 'cuda' if torch.cuda.is_available() else 'cpu' +global print +from MagicDec.Engine.tp import init_dist +use_tp = len(args.rank_group) > 1 +global_group = None +if use_tp: + rank, global_group = init_dist() + if rank != args.rank_group[0]: + print = lambda *args, **kwargs: None + +setup_seed(args.seed) +print(f"Using device={DEVICE}") + +MAX_LEN_TARGET = args.max_len +DTYPE = torch.bfloat16 +BATCH_SIZE = args.B +benchmark = args.benchmark +checkpoint_path = args.model + +target_dec_len = args.gamma + 1 +draft_dec_len = 1 + +# Load target model +engine = LMBackend(dtype=DTYPE, device=DEVICE, dec_len=target_dec_len) +engine.load_model(checkpoint_path, use_tp=use_tp, rank_group = args.rank_group, group=global_group) +vocab_size = engine.model.config.vocab_size + +# load draft model +draft_engine = LMBackendDraft(dtype=DTYPE, device=DEVICE, dec_len=draft_dec_len) +draft_engine.load_model(args.draft, use_tp=use_tp, rank_group = args.rank_group, group=global_group) + +if args.compile: + engine.compile() + draft_engine.compile() + +# Setup caches +engine.setup_caches(max_batch_size=BATCH_SIZE, max_seq_length=MAX_LEN_TARGET) +draft_engine.setup_caches(max_batch_size=BATCH_SIZE, max_seq_length=MAX_LEN_TARGET) # no sparse attention as of now + +# Load dataset +tokenizer = AutoTokenizer.from_pretrained(args.model_name) +tokenizer.pad_token = tokenizer.eos_token +eot_1 = tokenizer.eos_token_id +if tokenizer.unk_token_id is not None: + eot_2 = tokenizer.unk_token_id +else: + eot_2 = tokenizer.encode("<|eot_id|>")[-1] +print(f"eot_1: {eot_1}, eot_2: {eot_2}") + +if args.dataset == "pg19": + dataset = convert_pg19_dataset(tokenizer=tokenizer, seq_len=args.prefix_len) +elif args.dataset.startswith("ruler"): + dataset = convert_ruler_dataset(tokenizer=tokenizer, task=args.dataset.split(":")[1], model_name=args.model_name, seq_len=args.prefix_len) +else: + raise ValueError(f"Unknown dataset {args.dataset}") +dataloader = DataLoader(dataset, batch_size=BATCH_SIZE, shuffle=False, drop_last=True) +num_eval_steps = min(50, len(dataloader)) + +total_time = 0.0 +num_gen_tokens = 0 +target_steps = 0 +if benchmark: + draft_time = 0.0 + target_time = 0.0 + verify_loop = 0.0 + +for step, batch in tqdm(enumerate(dataloader), total=num_eval_steps): + # if step < 26: + # continue + if step >= num_eval_steps: + break + input_ids = batch[0].to(DEVICE) + terminal = False + tokens_buffer= torch.zeros((BATCH_SIZE, args.gamma+1), device=DEVICE).long() + output = torch.zeros(BATCH_SIZE, MAX_LEN_TARGET+1, device=DEVICE).long() + output[:, :input_ids.shape[1]] = input_ids + num_nodes = torch.zeros(BATCH_SIZE,device=DEVICE).long() + num_nodes += input_ids.shape[1] + + tokens_buffer[:, :1] = engine.encode(input_ids)[:, -1:] + # draft prefill + draft_engine.encode(input_ids=input_ids) + + # TODO: simplify code + next_double = False + double_buffer = None + + torch.cuda.synchronize() + start = time.perf_counter() + while terminal == False: + + # Draft speculation + if benchmark: + torch.cuda.synchronize() + t1 = time.time() + + for i in range(args.gamma): + if i == 0 and next_double: + next_tokens = draft_engine.inference(double_buffer) + tokens_buffer[:, i+1:i+2] = next_tokens.gather(1, cachelens_update.view(-1,1) - 1) + draft_engine.cachelens = draft_engine.cachelens - (2 - cachelens_update) + next_double = False + else: + tokens_buffer[:, i+1:i+2] = draft_engine.inference(tokens_buffer[:, i:i+1]) + + if benchmark: + torch.cuda.synchronize() + t2 = time.time() + draft_time += t2-t1 + + # Target verification + target_tokens = engine.inference(tokens_buffer) + # target_tokens = [] + # for i in range(args.gamma+1): + # target_tokens.append(engine.inference(tokens_buffer[:, i:i+1])) + # target_tokens = torch.cat(target_tokens, dim=1) + + if benchmark: + torch.cuda.synchronize() + t3 = time.time() + target_time+=t3-t2 + + target_steps+=1 + + # Vectorized Verify Loop + draft_tokens = tokens_buffer[:, 1:args.gamma+1] + flag_accept_matrix = (target_tokens[:, :args.gamma] == draft_tokens) # shape: (BATCH_SIZE, gamma) + eot_condition = ((draft_tokens == eot_1) | (draft_tokens == eot_2)) # shape: (BATCH_SIZE, gamma) + + # Compute accept_flags by considering both the acceptance condition and EOT tokens + accept_flags_int = (flag_accept_matrix & (~eot_condition)).int() + accept_flags_cumprod = torch.cumprod(accept_flags_int, dim=1) + accept_flags_matrix = accept_flags_cumprod.bool() + + # Compute the number of accepted tokens + accept_nums = accept_flags_matrix.sum(dim=1, keepdim=True) + 1 # shape: (BATCH_SIZE, 1) + + # if (accept_nums != args.gamma + 1).any(): + # import pdb; pdb.set_trace() + + # Check for termination conditions + condition = (eot_condition & accept_flags_matrix).any(dim=1, keepdim=True) + if condition.any(): + terminal = True + + # Rollback the memory length + engine.cachelens = engine.cachelens - args.gamma - 1 + engine.paged_kv_last_page_len = engine.paged_kv_last_page_len - args.gamma - 1 + draft_engine.cachelens = draft_engine.cachelens - args.gamma -1 + draft_engine.paged_kv_last_page_len = draft_engine.paged_kv_last_page_len - args.gamma -1 + + # Put the accepted tokens to output + positions = torch.arange(output.shape[1], device=DEVICE).view(1, -1).repeat(BATCH_SIZE, 1) + mask = (positions < (engine.cachelens.view(-1,1) + accept_nums)) & (positions >= engine.cachelens.view(-1, 1)) + positions_buffer = torch.arange(args.gamma+1, device=DEVICE).view(1, -1).repeat(BATCH_SIZE, 1) + mask_buffer = positions_buffer MAX_LEN_TARGET: + terminal = True + # Put Bonus tokens to the tokens buffer, and prepare the variables for next itr + if not terminal: + tokens_buffer[:, :1] = bonus_tokens + if accept_nums.max() == args.gamma + 1: + next_double = True + double_buffer = torch.zeros((BATCH_SIZE, 2), device=DEVICE).long() + mask = (accept_nums == args.gamma + 1).squeeze() + double_buffer[:, 0] = torch.where(mask, tokens_buffer[:, -1], bonus_tokens[:, 0]) + double_buffer[:, 1] = torch.where(mask, bonus_tokens[:, 0], torch.zeros_like(bonus_tokens[:, 0])) + # non_zero_mask = double_buffer != 0 + non_zero_mask = torch.ones_like(double_buffer, dtype=torch.bool) + non_zero_mask[:, 1] = mask + cachelens_update = non_zero_mask.sum(dim=1).flatten() + + if not terminal: + if benchmark: + torch.cuda.synchronize() + t4 = time.time() + verify_loop += t4-t3 + else: + for i in range(BATCH_SIZE): + output[i, num_nodes[i]] = bonus_tokens[i] + num_nodes += 1 + if benchmark: + torch.cuda.synchronize() + t4 = time.time() + verify_loop += t4-t3 + + torch.cuda.synchronize() + end=time.perf_counter() + total_time += end-start + num_gen_tokens += (num_nodes.sum() - (input_ids.shape[1]+1)*BATCH_SIZE) + if args.printoutput: + for i in range(BATCH_SIZE): + print(tokenizer.decode(output[i, args.prefix_len:num_nodes[i]])) + print("total time :{:.5f}s, time per iter :{:.5f}s, decoding step: {}, large model step: {}".format(total_time, total_time / target_steps, num_gen_tokens, target_steps)) + if benchmark: + print("target time :{:.5f}s, draft time :{:.5f}s, verify loop : {}, avg generate len per sentence: {}".format(target_time/target_steps, draft_time / target_steps, verify_loop/target_steps, num_gen_tokens/target_steps/BATCH_SIZE)) + if step < 3: # TODO: revert to 10? + total_time = 0.0 + num_gen_tokens = 0 + target_steps = 0 + if benchmark: + draft_time = 0.0 + target_time = 0.0 + verify_loop = 0.0 + if use_tp: + dist.barrier() \ No newline at end of file