Download code/models/tt_transformers/tt/common.py from tt-hous/clef: direct link, hf CLI and curl.
- Browser
- Download file 41.8 kB
-
https://huggingface.co/tt-hous/clef/resolve/main/code/models/tt_transformers/tt/common.py
- Command line
-
hf download hf://tt-hous/clef/code/models/tt_transformers/tt/common.py
-
curl -L -o common.py https://huggingface.co/tt-hous/clef/resolve/main/code/models/tt_transformers/tt/common.py
41.8 kB
| # SPDX-FileCopyrightText: © 2024 Tenstorrent USA, Inc. | |
| # SPDX-License-Identifier: Apache-2.0 | |
| import math | |
| import os | |
| import re | |
| from enum import Enum | |
| from types import SimpleNamespace | |
| from typing import List, Optional, Union | |
| import torch | |
| from loguru import logger | |
| from PIL import Image as PIL_Image | |
| from pydantic import AliasChoices, BaseModel, ConfigDict, Field | |
| import ttnn | |
| from models.common.tensor_utils import get_rot_transformation_mat as get_rot_transformation_mat_v2 | |
| class URL(BaseModel): | |
| uri: str | |
| def __str__(self) -> str: | |
| return self.uri | |
| class ImageMedia(BaseModel): | |
| image: Union[PIL_Image.Image, URL] | |
| model_config = ConfigDict(arbitrary_types_allowed=True) | |
| class Role(Enum): | |
| system = "system" | |
| user = "user" | |
| assistant = "assistant" | |
| ipython = "ipython" | |
| InterleavedTextMedia = Union[ | |
| str, | |
| # Specific modalities can be placed here, but not generic attachments | |
| # since models don't consume them in a generic way | |
| ImageMedia, | |
| List[Union[str, ImageMedia]], | |
| ] | |
| class Mode(Enum): | |
| DECODE = "decode" | |
| PREFILL = "prefill" | |
| class HostEmbedding(torch.nn.Module): | |
| def __init__(self, model_args): | |
| super().__init__() | |
| self.emb = torch.nn.Embedding(model_args.vocab_size, model_args.dim) | |
| def forward(self, x): | |
| return self.emb(x) | |
| class HostScaledEmbedding(HostEmbedding): | |
| def __init__(self, model_args): | |
| super().__init__(model_args) | |
| self.embed_scale = model_args.embed_scale | |
| def forward(self, x): | |
| return self.emb(x) * self.embed_scale | |
| # Default configuration for Paged Attention | |
| class PagedAttentionConfig: | |
| def __init__(self, block_size=32, max_num_blocks=1024): | |
| self.block_size = block_size | |
| self.max_num_blocks = max_num_blocks | |
| class RopeScalingType(str, Enum): | |
| """Types of RoPE scaling.""" | |
| # DYNAMIC = "dynamic" | |
| LINEAR = "linear" | |
| YARN = "yarn" | |
| LLAMA3 = "llama3" | |
| PHI3 = "longrope" | |
| DEFAULT = "default" | |
| class RopeScaling(BaseModel): | |
| """RoPE scaling configuration.""" | |
| rope_type: RopeScalingType = Field( | |
| validation_alias=AliasChoices("rope_type", "type"), exclude=True, description="RoPE scaling type" | |
| ) | |
| factor: Optional[float] = None | |
| original_max_position_embeddings: Optional[int] = None | |
| class RopeScalingLinear(RopeScaling): | |
| """RoPE scaling configuration for linear.""" | |
| class RopeScalingLlama3(RopeScaling): | |
| """RoPE scaling configuration for Llama-3.x.""" | |
| # Llama-3.x specific parameters | |
| low_freq_factor: Optional[float] = 1.0 | |
| high_freq_factor: Optional[float] = 4.0 | |
| class RopeScalingYarn(RopeScaling): | |
| """RoPE scaling configuration for Yarn.""" | |
| # Yarn-specific parameters | |
| beta_fast: Optional[float] = 32.0 | |
| beta_slow: Optional[float] = 1.0 | |
| mscale: Optional[float] = 1.0 | |
| mscale_all_dim: Optional[float] = 0.0 | |
| truncate: Optional[bool] = True # Whether to truncate the correction range (floor/ceil) | |
| class RopeScalingPhi3(RopeScaling): | |
| """RoPE scaling configuration for Phi3.""" | |
| # Phi3-specific parameters | |
| long_factor: Optional[list] | |
| short_factor: Optional[list] | |
| def rope_scaling_model_factory( | |
| rope_scaling_params: dict, original_max_context_len: Optional[int] = None | |
| ) -> RopeScaling: | |
| rope_scaling_type = rope_scaling_params.get("rope_type") or rope_scaling_params.get("type") | |
| if rope_scaling_type == RopeScalingType.LINEAR: | |
| return RopeScalingLinear(**rope_scaling_params) | |
| elif rope_scaling_type == RopeScalingType.LLAMA3: | |
| return RopeScalingLlama3(**rope_scaling_params) | |
| elif rope_scaling_type == RopeScalingType.YARN: | |
| return RopeScalingYarn(**rope_scaling_params) | |
| elif rope_scaling_type == RopeScalingType.PHI3: | |
| # transformers 5.x includes original_max_position_embeddings in the rope dict, | |
| # which collides with the explicit kwarg; merge so the caller value wins and the | |
| # key is only passed once. | |
| phi3_params = dict(rope_scaling_params) | |
| if original_max_context_len is not None: | |
| phi3_params["original_max_position_embeddings"] = original_max_context_len | |
| return RopeScalingPhi3(**phi3_params) | |
| elif rope_scaling_type in ["default", "mrope"]: | |
| logger.warning( | |
| f"Rope scaling type was set to {rope_scaling_type}, defaulting to no rope scaling as this rope type is not supported yet by TTT" | |
| ) | |
| return None | |
| else: | |
| raise ValueError(f"Unexpected RoPE scaling type: {rope_scaling_type}") | |
| # transformers 5.x consolidated the RoPE config: the top-level `rope_theta` / | |
| # `rope_local_base_freq` / `rope_scaling` keys were replaced by a single nested | |
| # `rope_parameters` dict (flat for Qwen/Llama; per-attention-type sub-dicts — | |
| # `full_attention` / `sliding_attention` — for Gemma-style models). The helpers | |
| # below read from either layout so configs from transformers <5 and >=5 work. | |
| def get_rope_theta(config: dict, default=None): | |
| """RoPE base period (global / full-attention).""" | |
| if config.get("rope_theta") is not None: | |
| return config["rope_theta"] | |
| rope_parameters = config.get("rope_parameters") or {} | |
| if rope_parameters.get("rope_theta") is not None: # flat (Qwen/Llama) | |
| return rope_parameters["rope_theta"] | |
| return (rope_parameters.get("full_attention") or {}).get("rope_theta", default) # Gemma-style | |
| def get_rope_local_base_freq(config: dict, default=None): | |
| """Gemma sliding-window local RoPE base (was top-level `rope_local_base_freq`).""" | |
| if config.get("rope_local_base_freq") is not None: | |
| return config["rope_local_base_freq"] | |
| rope_parameters = config.get("rope_parameters") or {} | |
| return (rope_parameters.get("sliding_attention") or {}).get("rope_theta", default) | |
| def get_rope_scaling(config: dict): | |
| """RoPE scaling params (factor, original_max_position_embeddings, rope_type, ...). | |
| transformers <5 put these under `rope_scaling`; >=5 merges them into | |
| `rope_parameters` (flat, or `full_attention` for Gemma-style). Returns the | |
| holding dict, or None when no non-default scaling is configured. | |
| """ | |
| rope_scaling = config.get("rope_scaling") | |
| if rope_scaling: | |
| return rope_scaling | |
| rope_parameters = config.get("rope_parameters") or {} | |
| if "full_attention" in rope_parameters: # Gemma-style nesting | |
| rope_parameters = rope_parameters.get("full_attention") or {} | |
| # Only a non-default rope_type carries scaling (factor, etc.). | |
| if rope_parameters.get("rope_type") not in (None, "default"): | |
| return rope_parameters | |
| return None | |
| # Minimal addition for Mistral vision support | |
| def position_ids_in_meshgrid_tt(tt_patch_embeds_list, max_width, device): | |
| position_ids_tt = [] | |
| for tt_patch in tt_patch_embeds_list: | |
| shape = tt_patch.shape | |
| height, width = shape[-2], shape[-1] | |
| mesh = torch.meshgrid(torch.arange(height), torch.arange(width), indexing="ij") | |
| h_grid, v_grid = torch.stack(mesh, dim=-1).reshape(-1, 2).chunk(2, -1) | |
| ids = h_grid * max_width + v_grid | |
| tt_ids = ttnn.from_torch( | |
| ids, | |
| device=device, | |
| dtype=ttnn.uint32, | |
| layout=ttnn.ROW_MAJOR_LAYOUT, | |
| memory_config=ttnn.DRAM_MEMORY_CONFIG, | |
| ) | |
| position_ids_tt.append(tt_ids[:, 0]) | |
| return ttnn.concat(position_ids_tt, dim=0) | |
| def encode_prompt_instruct(tokenizer, prompt_text, system_prompt_text=None): | |
| """<|begin_of_text|><|start_header_id|>system<|end_header_id|> | |
| {{ system_prompt }}<|eot_id|><|start_header_id|>user<|end_header_id|> | |
| {{ user_msg_1 }}<|eot_id|><|start_header_id|>assistant<|end_header_id|> | |
| {{ model_answer_1 }}<|eot_id|> | |
| """ | |
| begin_of_text = [tokenizer.special_tokens["<|begin_of_text|>"]] | |
| start_header = [tokenizer.special_tokens["<|start_header_id|>"]] | |
| end_header = [tokenizer.special_tokens["<|end_header_id|>"]] | |
| end_turn = [tokenizer.special_tokens["<|eot_id|>"]] | |
| system = tokenizer.encode("system", bos=False, eos=False) | |
| user = tokenizer.encode("user", bos=False, eos=False) | |
| assistant = tokenizer.encode("assistant", bos=False, eos=False) | |
| prompt = tokenizer.encode(prompt_text, bos=False, eos=False) | |
| system_prompt = start_header + system + end_header + system_prompt_text + end_turn if system_prompt_text else [] | |
| user_prompt = start_header + user + end_header + prompt + end_turn | |
| assistant_reply = start_header + assistant + end_header | |
| return begin_of_text + system_prompt + user_prompt + assistant_reply | |
| def preprocess_inputs_prefill( | |
| input_prompts, | |
| tokenizer, | |
| model_args, | |
| instruct, | |
| max_generated_tokens, | |
| max_prefill_len=128 * 1024, | |
| ): | |
| """ | |
| Run tokenizer on inputs, and create embeddings for the first token of each input | |
| """ | |
| # To avoid going out of memory, clip the max prefill length by the maximum number of tokens that will be generated | |
| for m_args in model_args: | |
| assert ( | |
| max_prefill_len <= m_args.max_context_len | |
| ), f"max_prefill_len {max_prefill_len} cannot exceed max_context_len {m_args.max_context_len}" | |
| # we need to make room for the generated tokens in the total token budget | |
| max_prefill_len -= max_generated_tokens | |
| assert ( | |
| max_prefill_len > 0 | |
| ), f"max_prefill_len ({max_prefill_len + max_generated_tokens}) must be greater than max_generated_tokens ({max_generated_tokens})" | |
| encoded_prompts = [ | |
| model_args[idx % len(model_args)].encode_prompt(prompt, instruct=instruct) | |
| for idx, prompt in enumerate(input_prompts) | |
| ] | |
| # Print the length of encoded prompts | |
| logger.info("Encoded prompt lengths:" + ", ".join(str(len(prompt)) for prompt in encoded_prompts)) | |
| prompt_lens = [len(x) for x in encoded_prompts] | |
| min_prompt_len = min(prompt_lens) | |
| max_prompt_len = max(prompt_lens) | |
| # To avoid running out of memory when giving prompts larger than the maximum, clip to max_prefill_len | |
| if min_prompt_len > max_prefill_len: | |
| logger.info(f"Left-clipping prompts to {max_prefill_len}") | |
| if instruct: | |
| # We need to allow a few tokens for the system prompt and the special turn tokens for assistant and user; | |
| # to find out how big those will be, we will: | |
| # 1. Tokenize the entire prompt with non-instruct tokenization | |
| # 2. Calculate overhead = length of instruct tokenization - length of non-instruct tokenization | |
| # 3. Shorten the tokenized clipped prompt by the overhead and convert back to text | |
| # 4. Tokenize the result with instruct tokenization | |
| # 5. Assert that the length of this is equal to the max_prefill_len | |
| raw_prompts = [ | |
| model_args[idx % len(model_args)].encode_prompt(prompt, instruct=False) | |
| for idx, prompt in enumerate(input_prompts) | |
| ] | |
| overhead = [len(e) - len(r) for e, r in zip(encoded_prompts, raw_prompts)] | |
| shortened = [] | |
| for idx, (e, o) in enumerate(zip(raw_prompts, overhead)): | |
| if isinstance(tokenizer, list): | |
| sp = tokenizer[idx % len(model_args)].decode(e[-(max_prefill_len - o) :]) | |
| else: | |
| sp = tokenizer.decode(e[-(max_prefill_len - o) :]) | |
| shortened.append(sp) | |
| encoded_prompts = [ | |
| model_args[idx % len(model_args)].encode_prompt(prompt, instruct=instruct) | |
| for idx, prompt in enumerate(shortened) | |
| ] | |
| # Instruct re-tokenization can drift by a few tokens vs the overhead | |
| # estimate (seen on Gemma4-26B-A4B: 65337 vs 65336). Re-trim / accept | |
| # slightly-short prompts rather than hard-failing the demo. | |
| trimmed = [] | |
| for e in encoded_prompts: | |
| if len(e) > max_prefill_len: | |
| e = e[-max_prefill_len:] | |
| trimmed.append(e) | |
| encoded_prompts = trimmed | |
| lens = [len(e) for e in encoded_prompts] | |
| assert all( | |
| 0 < n <= max_prefill_len for n in lens | |
| ), f"Clipped prompts are not of the correct length, expected <= {max_prefill_len} but got {lens}" | |
| if any(n != max_prefill_len for n in lens): | |
| logger.warning( | |
| f"Instruct re-clip lengths {lens} != target {max_prefill_len}; " | |
| f"continuing with trimmed/short prompts" | |
| ) | |
| else: | |
| encoded_prompts = [encod[-max_prefill_len:] for encod in encoded_prompts] | |
| # Update prompt lengths | |
| prompt_lens = [len(x) for x in encoded_prompts] | |
| min_prompt_len = min(prompt_lens) | |
| max_prompt_len = max(prompt_lens) | |
| for m in model_args: | |
| assert ( | |
| max_prompt_len <= m.max_seq_len | |
| ), f"Max prompt length {max_prompt_len} exceeds model max seq len {m.max_seq_len}" | |
| assert min_prompt_len > 0, "Minimum prompt length must be greater than 0" | |
| assert min_prompt_len <= max_prompt_len, f"Minimum prompt length {min_prompt_len} exceeds max len {max_prompt_len}" | |
| logger.info(f"# of users: {len(encoded_prompts)}") | |
| input_tokens_prefill = [] | |
| decoding_pos = [] | |
| prefill_lens = [] | |
| # Pad each prompt to the maximum length among all prompts. | |
| # To avoid issues, we keep track of the decoding position to decode correctly the user's prompt | |
| for i, encoded in enumerate(encoded_prompts): | |
| # Initial prefill tensors full of pad tokens | |
| input_tokens_prefill_i = torch.full((1, max_prompt_len), 0, dtype=torch.int32) | |
| input_tokens_prefill_i[0, : len(encoded[:])] = torch.tensor(encoded[:]).to(input_tokens_prefill_i) | |
| input_tokens_prefill.append(input_tokens_prefill_i) | |
| # Keep the correct decoding position of each user | |
| decoding_pos.append(len(encoded)) | |
| prefill_lens.append(max_prompt_len) | |
| return ( | |
| input_tokens_prefill, | |
| encoded_prompts, | |
| decoding_pos, | |
| prefill_lens, | |
| ) | |
| def _chat_template_ids(encoded): | |
| """Normalize apply_chat_template(tokenize=True) output to a flat List[int]. | |
| transformers <5 returned a plain List[int]; transformers 5.x defaults | |
| apply_chat_template to ``return_dict=True`` and returns a ``BatchEncoding`` | |
| (a ``UserDict`` — NOT a ``dict`` subclass, so ``isinstance(x, dict)`` is | |
| False), or a `tokenizers.Encoding` (exposes ``.ids``). Iterating a | |
| ``BatchEncoding``/``UserDict`` yields its *keys* ("input_ids", ...), so we | |
| must extract ``input_ids`` via mapping membership rather than ``isinstance``. | |
| """ | |
| # dict / BatchEncoding / UserDict — use mapping membership, since BatchEncoding | |
| # is a UserDict and fails isinstance(x, dict). | |
| if hasattr(encoded, "keys") and "input_ids" in encoded: | |
| encoded = encoded["input_ids"] | |
| if hasattr(encoded, "ids"): # tokenizers.Encoding | |
| return list(encoded.ids) | |
| if hasattr(encoded, "tolist"): # torch tensor / np array | |
| encoded = encoded.tolist() | |
| # apply_chat_template(return_dict=True) on a single conversation can nest the | |
| # ids in a 1-element batch dim ([[ids]]); unwrap it. | |
| if isinstance(encoded, (list, tuple)) and len(encoded) == 1 and isinstance(encoded[0], (list, tuple)): | |
| encoded = encoded[0] | |
| return list(encoded) # already a List[int] | |
| def encode_prompt_hf(tokenizer, prompt_text, system_prompt_text=None): | |
| """See https://huggingface.co/docs/transformers/main/en/chat_templating""" | |
| chat = [] | |
| if isinstance(prompt_text, str): | |
| if system_prompt_text: | |
| chat.append({"role": "system", "content": system_prompt_text}) | |
| if prompt_text: | |
| chat.append({"role": "user", "content": prompt_text}) | |
| encoded = tokenizer.apply_chat_template(chat, add_generation_prompt=True, tokenize=True) | |
| else: | |
| encoded = tokenizer.apply_chat_template(prompt_text, add_generation_prompt=True, tokenize=True) | |
| return _chat_template_ids(encoded) | |
| def compute_llama3_parameters(freqs: torch.Tensor, scale_factor: float, orig_context_len: int): | |
| """Llama-3.x specific scaling for rotary embeddings.""" | |
| low_freq_factor = 1 | |
| high_freq_factor = 4 | |
| low_freq_wavelen = orig_context_len / low_freq_factor | |
| high_freq_wavelen = orig_context_len / high_freq_factor | |
| new_freqs = [] | |
| for freq in freqs: | |
| wavelen = 2 * math.pi / freq | |
| if wavelen < high_freq_wavelen: | |
| new_freqs.append(freq) | |
| elif wavelen > low_freq_wavelen: | |
| new_freqs.append(freq / scale_factor) | |
| else: | |
| assert low_freq_wavelen != high_freq_wavelen | |
| smooth = (orig_context_len / wavelen - low_freq_factor) / (high_freq_factor - low_freq_factor) | |
| new_freqs.append((1 - smooth) * freq / scale_factor + smooth * freq) | |
| return torch.tensor(new_freqs, dtype=freqs.dtype, device=freqs.device) | |
| def compute_linear_parameters(freqs: torch.Tensor, scale_factor: float, orig_context_len: int): | |
| """Linear scaling for rotary embeddings.""" | |
| freqs /= scale_factor | |
| return freqs | |
| def compute_default_parameters(freqs: torch.Tensor, scale_factor: float, orig_context_len: int): | |
| """Default scaling for rotary embeddings.""" | |
| return freqs | |
| def apply_scaling(freqs: torch.Tensor, scale_factor: float, orig_context_len: int, rope_type="llama3"): | |
| # FIXME: Llama-3.x specific scaling - we need to support yarn for Qwen2.5 models | |
| if rope_type == "default": | |
| freqs = compute_default_parameters(freqs, scale_factor, orig_context_len) | |
| elif rope_type == "linear": | |
| freqs = compute_linear_parameters(freqs, scale_factor, orig_context_len) | |
| elif rope_type == "llama3": | |
| freqs = compute_llama3_parameters(freqs, scale_factor, orig_context_len) | |
| return freqs | |
| # Minimal addition for Mistral vision RoPE support | |
| def apply_scaling_vision(freqs: torch.Tensor, scale_factor: float, orig_context_len: int): | |
| return freqs / scale_factor | |
| # Minimal addition for Mistral vision RoPE support | |
| def precompute_mistral_vision_freqs( | |
| dim: int, max_patches_per_side: int, theta: float, scale_factor=None, orig_context_len=None | |
| ): | |
| # Compute base frequencies | |
| base_freqs = 1.0 / (theta ** (torch.arange(0, dim, 2).float() / dim)) | |
| if scale_factor is not None: | |
| base_freqs = apply_scaling_vision(base_freqs, scale_factor, orig_context_len) | |
| # Get height and width indices | |
| h_idx = torch.arange(max_patches_per_side) | |
| w_idx = torch.arange(max_patches_per_side) | |
| # Compute 2D frequency matrices | |
| freqs_h = torch.outer(h_idx, base_freqs[::2]) | |
| freqs_w = torch.outer(w_idx, base_freqs[1::2]) | |
| # Broadcast + merge | |
| inv_freq = torch.cat( | |
| [ | |
| freqs_h[:, None, :].repeat(1, max_patches_per_side, 1), | |
| freqs_w[None, :, :].repeat(max_patches_per_side, 1, 1), | |
| ], | |
| dim=-1, | |
| ).reshape( | |
| -1, dim // 2 | |
| ) # Shape: [H*W, dim//2] | |
| full_freqs = torch.cat([inv_freq, inv_freq], dim=-1) | |
| cos = full_freqs.cos() | |
| sin = full_freqs.sin() | |
| return cos, sin # Shape: [H*W, dim] | |
| def precompute_freqs(dim: int, end: int, theta, scale_factor, orig_context_len, rope_type="llama3"): | |
| """ | |
| Precompute the frequency tensor for sine and cosine values with given dimensions. | |
| Args: | |
| dim (int): Dimension of the frequency tensor. | |
| end (int): End index for precomputing frequencies. | |
| theta (float, optional): Scaling factor for frequency computation. Defaults to 500000.0. | |
| Returns: | |
| Tuple[torch.Tensor, torch.Tensor]: Tensors containing cosine and sine values. | |
| """ | |
| freqs = 1.0 / (theta ** (torch.arange(0, dim, 2)[: (dim // 2)].float() / dim)) | |
| t = torch.arange(end) | |
| if scale_factor is not None: | |
| freqs = apply_scaling(freqs, scale_factor, orig_context_len, rope_type=rope_type) | |
| freqs = torch.outer(t, freqs).float() | |
| return torch.cos(freqs), torch.sin(freqs) | |
| def freqs_to_rotation_matrix(cos_freqs, sin_freqs): | |
| """ | |
| Transform cos/sin frequencies to a rotation matrix. | |
| """ | |
| emb_size, emb_dim = cos_freqs.shape | |
| dhead = emb_dim * 2 | |
| rot_emb_matrix = torch.zeros(emb_size, dhead, dhead) | |
| rot_emb_matrix[..., torch.arange(0, dhead, 2), torch.arange(0, dhead, 2)] = cos_freqs.clone() | |
| rot_emb_matrix[..., torch.arange(1, dhead, 2), torch.arange(1, dhead, 2)] = cos_freqs.clone() | |
| rot_emb_matrix[..., torch.arange(0, dhead, 2), torch.arange(1, dhead, 2)] = -sin_freqs.clone() | |
| rot_emb_matrix[..., torch.arange(1, dhead, 2), torch.arange(0, dhead, 2)] = sin_freqs.clone() | |
| rot_emb_matrix = rot_emb_matrix.transpose(-1, -2) # Necessary for correct rotation when applied as (x @ R) | |
| return rot_emb_matrix | |
| def gather_cos_sin(position_ids, cos, sin): | |
| position_id_expanded = position_ids.unsqueeze(1).expand(-1, cos.shape[-1]) | |
| cos = cos.gather(0, position_id_expanded) | |
| sin = sin.gather(0, position_id_expanded) | |
| cos = torch.stack([cos, cos], dim=-1).flatten(-2).unsqueeze(0).unsqueeze(0) | |
| sin = torch.stack([sin, sin], dim=-1).flatten(-2).unsqueeze(0).unsqueeze(0) | |
| return cos, sin | |
| def get_prefill_rot_mat(head_dim, mesh_device, seq_len, theta, scale_factor, orig_context_len, start_pos=0): | |
| cos, sin = precompute_freqs( | |
| head_dim, seq_len * 2, theta=theta, scale_factor=scale_factor, orig_context_len=orig_context_len | |
| ) | |
| cos_gathered, sin_gathered = gather_cos_sin(torch.arange(start_pos, start_pos + seq_len), cos, sin) | |
| assert cos_gathered.size() == (1, 1, seq_len, head_dim) | |
| assert sin_gathered.size() == (1, 1, seq_len, head_dim) | |
| cos_gathereds = ttnn.from_torch( | |
| cos_gathered, | |
| dtype=ttnn.bfloat16, | |
| layout=ttnn.TILE_LAYOUT, | |
| device=mesh_device, | |
| mesh_mapper=ttnn.ReplicateTensorToMesh(mesh_device), | |
| ) | |
| sin_gathereds = ttnn.from_torch( | |
| sin_gathered, | |
| dtype=ttnn.bfloat16, | |
| layout=ttnn.TILE_LAYOUT, | |
| device=mesh_device, | |
| mesh_mapper=ttnn.ReplicateTensorToMesh(mesh_device), | |
| ) | |
| rot_mats = [cos_gathereds, sin_gathereds] | |
| return rot_mats | |
| # Add-Multiply method of rotary embeddings for prefill | |
| def get_rot_transformation_mat(dhead=32): | |
| # ROPE op uses a single tile | |
| dhead = 32 | |
| # Delegate to TTTv2 implementation for consistency | |
| return get_rot_transformation_mat_v2(dhead) | |
| def get_single_rot_mat( | |
| dhead, | |
| mesh_device, | |
| num_devices, | |
| start_pos, | |
| theta, | |
| scale_factor, | |
| orig_context_len, | |
| on_host=False, | |
| ): | |
| freqs_unscaled = 1.0 / (theta ** (torch.arange(0, dhead, 2)[: (dhead // 2)].float() / dhead)) | |
| if scale_factor is not None: | |
| freqs = apply_scaling(freqs_unscaled, scale_factor, orig_context_len, rope_type="llama3") | |
| rot_matrix = torch.zeros(dhead, dhead) | |
| # [INFO] freqs_unscaled and freqs are forced to float dtype above and it should be converted back to match dtype of rot_matrix | |
| sin_freqs, cos_freqs = torch.sin(freqs).to(rot_matrix.dtype), torch.cos(freqs).to(rot_matrix.dtype) | |
| rot_matrix[torch.arange(0, dhead, 2), torch.arange(0, dhead, 2)] = cos_freqs.clone() | |
| rot_matrix[torch.arange(1, dhead, 2), torch.arange(1, dhead, 2)] = cos_freqs.clone() | |
| rot_matrix[torch.arange(0, dhead, 2), torch.arange(1, dhead, 2)] = -sin_freqs.clone() | |
| rot_matrix[torch.arange(1, dhead, 2), torch.arange(0, dhead, 2)] = sin_freqs.clone() | |
| rot_matrix = rot_matrix.transpose(-1, -2) | |
| # Support for start_pos different than 0 | |
| freqs = start_pos * freqs_unscaled | |
| if scale_factor is not None: | |
| freqs = apply_scaling(freqs, scale_factor, orig_context_len, rope_type="llama3") | |
| current_rot_mat = torch.zeros(dhead, dhead) | |
| # [INFO] freqs_unscaled and freqs are forced to float dtype above and it should be converted back to match dtype of current_rot_mat | |
| sin_freqs, cos_freqs = torch.sin(freqs).to(current_rot_mat.dtype), torch.cos(freqs).to(current_rot_mat.dtype) | |
| current_rot_mat[torch.arange(0, dhead, 2), torch.arange(0, dhead, 2)] = cos_freqs.clone() | |
| current_rot_mat[torch.arange(1, dhead, 2), torch.arange(1, dhead, 2)] = cos_freqs.clone() | |
| current_rot_mat[torch.arange(0, dhead, 2), torch.arange(1, dhead, 2)] = -sin_freqs.clone() | |
| current_rot_mat[torch.arange(1, dhead, 2), torch.arange(0, dhead, 2)] = sin_freqs.clone() | |
| return ttnn.from_torch( | |
| current_rot_mat.T.unsqueeze(0).unsqueeze(0), # 1,1,head_dim,head_dim | |
| device=mesh_device if not on_host else None, | |
| dtype=ttnn.bfloat16, | |
| layout=ttnn.TILE_LAYOUT, | |
| mesh_mapper=ttnn.ReplicateTensorToMesh(mesh_device) if num_devices > 1 or not on_host else None, | |
| ), ttnn.from_torch( | |
| rot_matrix.unsqueeze(0).unsqueeze(0), # 1,1,head_dim,head_dim | |
| device=mesh_device if not on_host else None, | |
| dtype=ttnn.bfloat16, | |
| layout=ttnn.TILE_LAYOUT, | |
| mesh_mapper=ttnn.ReplicateTensorToMesh(mesh_device) if num_devices > 1 or not on_host else None, | |
| ) | |
| def num_to_core_range_set(x): | |
| assert x < 8 or x % 8 == 0 | |
| num_x = min(x, 8) | |
| num_y = x // num_x | |
| assert num_x * num_y == x | |
| return ttnn.CoreRangeSet( | |
| { | |
| ttnn.CoreRange( | |
| ttnn.CoreCoord(0, 0), | |
| ttnn.CoreCoord(num_x - 1, num_y - 1), | |
| ), | |
| } | |
| ) | |
| def copy_host_to_device( | |
| host_tensors, | |
| device_tensors=None, | |
| mesh_device=None, | |
| shard_specs=None, | |
| ): | |
| """ | |
| Helper function which copies host tensors to device tensors. | |
| If no device_tensors are provided, it creates new device tensors and returns them. | |
| """ | |
| if device_tensors is None: | |
| assert mesh_device is not None, "mesh_device is required when device_tensors is None" | |
| ret = [] | |
| for i in range(len(host_tensors)): | |
| if shard_specs and shard_specs[i] is not None: | |
| on_device = host_tensors[i].to(mesh_device, shard_specs[i]) if host_tensors[i] else None | |
| else: | |
| on_device = ttnn.to_device(host_tensors[i], device=mesh_device) if host_tensors[i] else None | |
| ret.append(on_device) | |
| return ret | |
| else: | |
| for i in range(len(host_tensors)): | |
| if host_tensors[i] is None: | |
| assert device_tensors[i] is None | |
| continue | |
| ttnn.copy_host_to_device_tensor(host_tensors[i], device_tensors[i]) | |
| return device_tensors | |
| def calculate_hidden_dim(dim, ffn_dim_multiplier, multiple_of): | |
| """Helper function based on logic used in reference model: | |
| https://github.com/meta-llama/llama-models/blob/e4a6ed52a142bb9b5106dcbf48e41f97f8e7378e/models/llama3/reference_impl/model.py#L227C7-L231C83 | |
| """ | |
| hidden_dim = int(2 * (4 * dim) / 3) | |
| if ffn_dim_multiplier is not None: | |
| hidden_dim = int(ffn_dim_multiplier * hidden_dim) | |
| hidden_dim = multiple_of * ((hidden_dim + multiple_of - 1) // multiple_of) | |
| return hidden_dim | |
| def get_out_subblock_w(per_core_N, out_subblock_h): | |
| """ | |
| Helper function to calculate the out_subblock_w based on the per_core_N and out_subblock_h | |
| """ | |
| out_subblock_w = 4 # TODO: Check with LLK team if this is the true bound, might be 8 now | |
| while out_subblock_w > 1: | |
| if out_subblock_w * out_subblock_h <= 4 and per_core_N % out_subblock_w == 0: | |
| break | |
| out_subblock_w -= 1 | |
| return out_subblock_w | |
| def first_five(tensor, mesh_device, start=0, end=5): | |
| """ | |
| Helper function to return the first 5 elements of a tensor via torch, or optionally another slice | |
| """ | |
| return torch.Tensor(ttnn.to_torch(tensor, mesh_composer=ttnn.ConcatMeshToTensor(mesh_device, dim=-1)))[ | |
| 0, 0, 0, start:end | |
| ] | |
| def last_five(tensor, mesh_device): | |
| """ | |
| Helper function to return the last 5 elements of a tensor via torch | |
| """ | |
| return torch.Tensor(ttnn.to_torch(tensor, mesh_composer=ttnn.ConcatMeshToTensor(mesh_device, dim=-1)))[0, 0, 0, -5:] | |
| # Sample logits from a distribution | |
| def sample_top_p(probs: torch.Tensor, p: float): | |
| assert 0 <= p <= 1 | |
| probs_sort, probs_idx = torch.sort(probs, dim=-1, descending=True) | |
| probs_sum = torch.cumsum(probs_sort, dim=-1) | |
| mask = probs_sum - probs_sort > p | |
| probs_sort[mask] = 0.0 | |
| probs_sort.div_(probs_sort.sum(dim=-1, keepdim=True)) | |
| next_token = torch.multinomial(probs_sort, num_samples=1) | |
| return torch.gather(probs_idx, -1, next_token) | |
| def sample_host(tt_input, temperature=0.6, top_p=0.08, on_host=True): | |
| vocab_size = tt_input.shape[-1] | |
| pt_input = tt_input[..., :vocab_size] | |
| if temperature > 0: | |
| probs = torch.softmax(pt_input / temperature, dim=-1) | |
| pt_out = sample_top_p(probs.squeeze(), top_p) | |
| else: | |
| pt_out = torch.argmax(pt_input, dim=-1) | |
| if pt_out.dim() == 1: # if sampling a single token re-add the batch dim to the tensor | |
| pt_out = pt_out.unsqueeze(0) | |
| return None, pt_out | |
| def get_padded_prefill_len(seq_len: int) -> int: | |
| """ | |
| Get the padded prefill length for a given sequence length. | |
| This is used to pad the sequence length to the nearest power of 2. | |
| """ | |
| # TODO: https://github.com/tenstorrent/tt-metal/issues/34117 | |
| if seq_len <= 128: | |
| return 128 | |
| if seq_len <= 1024: | |
| return 1024 | |
| else: | |
| # return next power of 2 greater than seq_len | |
| return 2 ** (seq_len - 1).bit_length() | |
| def get_all_padded_prefill_lengths(max_len): | |
| lengths = [128] | |
| k = 0 | |
| while (v := (1 << k) * 1024) <= max_len: | |
| lengths.append(v) | |
| k += 1 | |
| return lengths | |
| def calculate_prefill_warmup_seq_lens(max_seq_len_to_warmup, trace_supported_seq_lens): | |
| to_warmup_seq_lens = get_all_padded_prefill_lengths(max_seq_len_to_warmup) | |
| for trace_supported_seq_len in trace_supported_seq_lens: | |
| if trace_supported_seq_len not in to_warmup_seq_lens: | |
| to_warmup_seq_lens.append(trace_supported_seq_len) | |
| to_warmup_seq_lens.sort() | |
| return to_warmup_seq_lens | |
| def cap_seq_lens_to_max_prefill_chunk_size(seq_lens, cap): | |
| for seq_len in seq_lens: | |
| if seq_len > cap: | |
| seq_lens = seq_lens[: seq_lens.index(seq_len)] | |
| break | |
| return seq_lens | |
| def get_block_size(kv_cache): | |
| return kv_cache[0][0].shape[2] | |
| def num_blocks_in_seq(seq_len, block_size): | |
| return math.ceil(seq_len / block_size) | |
| def nearest_pow_2(x): | |
| return 2 ** math.ceil(math.log2(x)) | |
| def get_max_prefill_chunk_size(seq_len, max_prefill_seq_len): | |
| """ | |
| Determine the largest multiple of 2048 that divides `seq_len` and is less than or equal to `max_prefill_seq_len`. | |
| **Assumptions**: | |
| - `seq_len` is a multiple of 2048. | |
| - `max_prefill_seq_len` is a multiple of 2048. | |
| """ | |
| MIN_CHUNK_SIZE = 2048 | |
| if not isinstance(seq_len, int) or not isinstance(max_prefill_seq_len, int): | |
| raise TypeError("Both seq_len and max_prefill_seq_len must be integers.") | |
| if seq_len <= 0 or max_prefill_seq_len <= 0: | |
| raise ValueError("Both seq_len and max_prefill_seq_len must be positive integers.") | |
| if seq_len % MIN_CHUNK_SIZE != 0: | |
| raise ValueError(f"seq_len ({seq_len}) must be a multiple of {MIN_CHUNK_SIZE}.") | |
| if max_prefill_seq_len % MIN_CHUNK_SIZE != 0: | |
| raise ValueError(f"max_prefill_seq_len ({max_prefill_seq_len}) must be a multiple of {MIN_CHUNK_SIZE}.") | |
| # Calculate the maximum possible chunk size | |
| # It cannot exceed either max_prefill_seq_len or seq_len | |
| max_possible_chunk = min(max_prefill_seq_len, seq_len) | |
| # Iterate from the largest possible multiple of MIN_CHUNK_SIZE down to MIN_CHUNK_SIZE | |
| for chunk_size in range(max_possible_chunk, 0, -MIN_CHUNK_SIZE): | |
| if seq_len % chunk_size == 0: | |
| return chunk_size | |
| raise ValueError("No valid chunk size found") | |
| def nearest_multiple(x, multiple_of): | |
| return math.ceil(x / multiple_of) * multiple_of | |
| def pad_to_size(x: torch.Tensor, dim: int, size: int) -> torch.Tensor: | |
| """ | |
| Pads the specified dimension of the input tensor with zeros | |
| :param x: Input PyTorch Tensor | |
| :param dim: The dimension to pad | |
| :param size: The size to pad to | |
| :return: Padded PyTorch Tensor | |
| """ | |
| # handle negative dim | |
| if dim < 0: | |
| dim = x.dim() + dim | |
| assert isinstance(x, torch.Tensor), "Input must be a torch.Tensor" | |
| assert -x.dim() <= dim < x.dim(), f"Dimension {dim} out of range (expected between {-x.dim()} and {x.dim() - 1})" | |
| dim = x.dim() + dim if dim < 0 else dim | |
| current_size = x.size(dim) | |
| pad_size = size - current_size | |
| if pad_size == 0: | |
| return x # No padding needed | |
| # Prepare the padding configuration for F.pad | |
| # F.pad expects padding in the form (pad_last_dim_left, pad_last_dim_right, ..., pad_dim_left, pad_dim_right) | |
| # We only pad on the "end" side of the specified dimension | |
| pad = [0] * (2 * x.dim()) # Initialize padding for all dimensions | |
| pad_index = 2 * (x.dim() - dim - 1) | |
| pad[pad_index + 1] = pad_size # Pad on the "right" side of the specified dimension | |
| padded_x = torch.nn.functional.pad(x, pad, mode="constant", value=0) | |
| return padded_x | |
| def get_base_model_name(model_name: str) -> str: | |
| # Explicitly handle phi-4 which doesn't follow the <Size>B format | |
| if "phi-4" in model_name.lower(): | |
| return "Phi-4" | |
| # Remove the suffix after B- (case insensitive), e.g. "Llama-3.1-70B-Instruct" -> "Llama-3.1-70B" | |
| match = re.search(r"(.*?\d+[bB])-", model_name) | |
| return match.group(1) if match else model_name | |
| def get_hf_model_name(model_path: str) -> str: | |
| # HF model name | |
| if model_path.count("/") == 1: | |
| return model_path | |
| # HF cache path | |
| pattern = r".*/?models--(?P<model_provider>[^/]+?)--(?P<model_name>[^/]+)/?" | |
| match = pattern.search(pattern, model_path) | |
| if match: | |
| model_provider = match.group("model_provider") | |
| model_name = match.group("model_name") | |
| return f"{model_provider}/{model_name}" | |
| raise ValueError( | |
| f"Unsupported '{model_path}', please use HF model name or follow HF format with 'models--<model_provider>--<model_name>'" | |
| ) | |
| def get_hf_tt_cache_path(model_path: str) -> str: | |
| tt_cache_home = os.getenv("TT_CACHE_HOME", "/mnt/MLPerf/huggingface/tt_cache/") | |
| if not os.path.exists(tt_cache_home): | |
| tt_cache_home = "model_cache" | |
| model_name = get_hf_model_name(model_path) | |
| tt_cache_path = os.path.join(tt_cache_home, model_name) | |
| if not os.path.exists(tt_cache_path): | |
| os.makedirs(tt_cache_path, exist_ok=True) | |
| return tt_cache_path | |
| def create_tt_model( | |
| mesh_device, | |
| instruct, | |
| max_batch_size, | |
| optimizations, | |
| max_seq_len, | |
| paged_attention_config: PagedAttentionConfig = None, | |
| dtype=ttnn.bfloat8_b, | |
| state_dict=None, | |
| num_layers=None, | |
| use_prefetcher=False, | |
| use_hf_rope=False, | |
| ): | |
| from models.tt_transformers.tt.model import Transformer | |
| from models.tt_transformers.tt.model_config import ModelArgs | |
| from models.tt_transformers.tt.prefetcher import Prefetcher | |
| num_tensors = 5 if use_prefetcher else 0 | |
| prefetcher = Prefetcher(mesh_device, num_tensors, num_layers) if use_prefetcher else None | |
| tt_model_args = ModelArgs( | |
| mesh_device, | |
| instruct=instruct, | |
| max_batch_size=max_batch_size, | |
| optimizations=optimizations, | |
| max_seq_len=max_seq_len, | |
| prefetcher=prefetcher, | |
| use_hf_rope=use_hf_rope, | |
| ) | |
| if num_layers is not None: | |
| tt_model_args.n_layers = num_layers | |
| if prefetcher is not None: | |
| prefetcher.num_layers = tt_model_args.n_layers | |
| # Decide whether the HF weights are still needed on host. When the ttnn weight cache for | |
| # this build was already fully built on a previous run, ttnn.as_tensor loads every weight from | |
| # disk and the state_dict is never read -- so skip the expensive from_pretrained host load | |
| # entirely (the load that OOMs/hangs in prefill, #48509). Generalizes GPT-OSS PR #48531 (whose | |
| # --skip-model-load pytest flag is gpt_oss-only; nothing equivalent exists for these models). | |
| # | |
| # state_dict is None -> decide here (warm cache => placeholder, else cold load). | |
| # state_dict falsy/{} -> caller already decided to skip (e.g. a prior DP submesh); build as-is. | |
| # state_dict populated -> reuse across DP models (avoid reloading for every submesh). | |
| loaded_real_weights = False | |
| if state_dict is None: | |
| if not tt_model_args.dummy_weights and tt_model_args.weight_cache_is_complete(dtype): | |
| logger.info("Warm ttnn weight cache detected -- skipping HF state_dict load.") | |
| # Dataless placeholder: every weight is loaded from its .tensorbin by ttnn.as_tensor; | |
| # the placeholder only satisfies the host-side reshape ops (see placeholder_state_dict). | |
| state_dict = tt_model_args.placeholder_state_dict(dtype) | |
| else: | |
| state_dict = tt_model_args.load_state_dict() | |
| loaded_real_weights = bool(state_dict) and not tt_model_args.dummy_weights | |
| # A populated state_dict handed in by the caller (DP submeshes after the first) bypasses | |
| # load_state_dict(), which is the only place the cold path sets is_mixture_of_experts. Without | |
| # this the later lanes build a dense MLP for an MoE checkpoint and fail on the missing | |
| # feed_forward.w1 key. Derive the flag from the keys, as load_state_dict does. | |
| # (The warm-cache placeholder mapping is deliberately falsy, so test for None, not truthiness.) | |
| if state_dict is not None and not getattr(tt_model_args, "is_mixture_of_experts", False): | |
| tt_model_args.is_mixture_of_experts = any(".experts." in k for k in state_dict.keys()) | |
| if getattr(tt_model_args, "is_mixture_of_experts", False): | |
| # Reused weights must initialize the same MoE configuration as load_state_dict. | |
| tt_model_args.moe = True | |
| expert_indices = [ | |
| int(k.split(".experts.")[1].split(".")[0]) + 1 for k in state_dict if "block_sparse_moe.experts." in k | |
| ] | |
| tt_model_args.num_experts = max(expert_indices) if expert_indices else tt_model_args.num_local_experts | |
| model = Transformer( | |
| args=tt_model_args, | |
| mesh_device=mesh_device, | |
| dtype=dtype, | |
| state_dict=state_dict, | |
| weight_cache_path=tt_model_args.weight_cache_path(dtype), | |
| paged_attention_config=paged_attention_config, | |
| prefetcher=prefetcher, | |
| ) | |
| # If this run populated the cache from a cold host load, record completion so future runs | |
| # can skip the load. Only for full-model builds (a num_layers override produces a partial | |
| # cache that must not satisfy the completeness check). | |
| if loaded_real_weights and num_layers is None: | |
| tt_model_args.mark_weight_cache_complete(dtype, state_dict) | |
| tt_kv_cache = [l.attention.layer_past for l in model.layers] if paged_attention_config else None | |
| return tt_model_args, model, tt_kv_cache, state_dict | |
| def hf_multimodal_encode(messages, processor): | |
| hf_messages = [] | |
| for msg in messages: | |
| hf_content = [] | |
| for item in msg.content: | |
| if isinstance(item, ImageMedia): | |
| hf_content.append( | |
| { | |
| "type": "image", | |
| "image": item.image, | |
| } | |
| ) | |
| elif isinstance(item, str): | |
| hf_content.append( | |
| { | |
| "type": "text", | |
| "text": item, | |
| } | |
| ) | |
| hf_messages.append( | |
| { | |
| "role": msg.role, | |
| "content": hf_content, | |
| } | |
| ) | |
| encoded = processor.apply_chat_template( | |
| hf_messages, add_generation_prompt=True, tokenize=True, return_dict=True, return_tensors="pt" | |
| ).to("cpu", dtype=torch.bfloat16) | |
| return SimpleNamespace( | |
| **encoded, | |
| tokens=encoded["input_ids"].squeeze(0), | |
| vision=SimpleNamespace( | |
| images=encoded.get("pixel_values", None), | |
| mask=None, | |
| ), | |
| ) | |
| def get_decode_mask(args, mesh_device, paged_attention_config=None): | |
| """Function to create a decoding mask for the attention mechanism.""" | |
| if paged_attention_config is not None: | |
| max_seq_len = (paged_attention_config.max_num_blocks * paged_attention_config.block_size) // args.max_batch_size | |
| else: | |
| max_seq_len = args.max_seq_len | |
| mask = torch.triu( | |
| torch.full( | |
| (args.max_batch_size, args.n_heads // mesh_device.shape[1], max_seq_len, max_seq_len), | |
| -float("inf"), | |
| dtype=torch.bfloat16, | |
| ), | |
| diagonal=1, | |
| ) | |
| if args.sliding_window > 0: | |
| mask += torch.tril( | |
| torch.full( | |
| (args.max_batch_size, args.n_heads // mesh_device.shape[1], max_seq_len, max_seq_len), | |
| -float("inf"), | |
| dtype=torch.bfloat16, | |
| ), | |
| diagonal=-args.sliding_window, | |
| ) | |
| return mask | |
| def build_encoder_attention_mask( | |
| x: torch.Tensor, | |
| ar: torch.Tensor, | |
| ntok: int, | |
| num_chunks: int, | |
| n_heads: int, | |
| ): | |
| """ | |
| Build vision encoder attention mask that omits padding tokens. | |
| """ | |
| def get_negative_inf_value(dtype): | |
| return torch.finfo(dtype).min | |
| masks = [] | |
| for arx in ar: | |
| mask_i = torch.ones((num_chunks, x.shape[2], 1), dtype=x.dtype) | |
| mask_i[: arx[0] * arx[1], :ntok] = 0 | |
| mask_i = mask_i.view(num_chunks * x.shape[2], -1) | |
| mask_i = mask_i @ mask_i.T * get_negative_inf_value(x.dtype) | |
| mask_i = mask_i.unsqueeze(0) | |
| masks.append(mask_i) | |
| masks = torch.stack(masks).to(x.device).expand(-1, n_heads, -1, -1) | |
| return masks | |