File size: 6,999 Bytes
be3ecc8 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 | # SPDX-FileCopyrightText: © 2025 Tenstorrent USA, Inc.
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass
from enum import Enum
from typing import Dict, List, Optional, Union
import torch
from PIL import Image
from pydantic import BaseModel, field_validator
# ``AutoModelForImageTextToText`` (the replacement for ``AutoModelForVision2Seq``,
# which was removed in transformers 5.x) is only consumed by the
# ``GeneratorChat``/``GeneratorText`` constructors below — defer the imports
# so loading this module doesn't break every downstream import chain
# (e.g. ``tt_transformers.tt.generator``, used by every TT vLLM bridge)
# under transformers >= 5.
class Role(Enum):
system = "system"
user = "user"
assistant = "assistant"
ipython = "ipython"
class StopReason(Enum):
end_of_turn = "end_of_turn"
end_of_message = "end_of_message"
out_of_tokens = "out_of_tokens"
@dataclass
class TokenResult:
token: int
text: str
logprobs: Optional[List[float]] = None
@dataclass
class CompletionMessage:
content: str
role: Role = Role.assistant.value
class BuiltinTool(Enum):
brave_search = "brave_search"
wolfram_alpha = "wolfram_alpha"
photogen = "photogen"
code_interpreter = "code_interpreter"
Primitive = Union[str, int, float, bool, None]
RecursiveType = Union[Primitive, List[Primitive], Dict[str, Primitive]]
class ToolCall(BaseModel):
call_id: str
tool_name: Union[BuiltinTool, str]
arguments: Dict[str, RecursiveType]
@field_validator("tool_name", mode="before")
@classmethod
def validate_field(cls, v):
if isinstance(v, str):
try:
return BuiltinTool(v)
except ValueError:
return v
return v
class ChatPrediction:
generation: CompletionMessage
decoded_tokens: Optional[List[str]] = None
logprobs: Optional[List[List[float]]] = None
class CompletionPrediction:
generation: str
decoded_tokens: Optional[List[str]] = None
logprobs: Optional[List[List[float]]] = None
def sample_top_p(probs, p):
"""
Perform top-p (nucleus) sampling on a probability distribution.
Args:
probs (torch.Tensor): Probability distribution tensor.
p (float): Probability threshold for top-p sampling.
Returns:
torch.Tensor: Sampled token indices.
Note:
Top-p sampling selects the smallest set of tokens whose cumulative probability mass
exceeds the threshold p. The distribution is renormalized based on the selected tokens.
From: https://github.com/meta-llama/llama-models/blob/v0.1.5/models/llama3/reference_impl/generation.py#L450-L472
"""
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)
next_token = torch.gather(probs_idx, -1, next_token)
return next_token
def extract_images_from_messages(messages):
images = []
for message in messages:
if "content" in message:
contents = message["content"]
for content in contents:
if (content["type"] == "image") and ("image" in content):
images.append(content["image"])
return images
def create_vision_mask(
tokens: List[int],
vision_token: int,
) -> List[List[int]]:
"""From: https://github.com/meta-llama/llama-models/blob/v0.1.5/models/llama3/api/chat_format.py#L253-L276"""
vision_token_locations = [i for i, token in enumerate(tokens) if token == vision_token]
if len(vision_token_locations) == 0:
return []
if len(vision_token_locations) == 1:
# only one image present, unmask until end of sequence
return [[vision_token_locations[0], -1]]
vision_masks = [[loc1, loc2] for loc1, loc2 in zip(vision_token_locations[:-1], vision_token_locations[1:])]
# last image will attend to all subsequent text
vision_masks.append([vision_token_locations[-1], len(tokens)])
# if there are two or more consecutive vision tokens,
# they should all attend to all subsequent
# text present
last_mask_end = vision_masks[-1][1]
for vision_mask in vision_masks[::-1]:
if vision_mask[0] == vision_mask[1] - 1:
vision_mask[1] = last_mask_end
last_mask_end = vision_mask[1]
return vision_masks
def encode_content(content, images, image_token):
if isinstance(content, Image):
images.append(content)
assert image_token is not None
return image_token
if isinstance(content, str):
return content
if isinstance(content, (list, tuple)):
return "\n".join(encode_content(item, images) for item in content)
if isinstance(content, dict):
content_type = content.get("type")
if content_type == "text":
return content["text"]
if content_type == "image":
# TBD: support url
images.append(content["image"])
assert image_token is not None
return image_token
raise ValueError(f"Unknown content format: {content}")
class GeneratorChat:
def __init__(self, model_name, max_batch_size=1):
from transformers import pipeline
self.pipe = pipeline("image-text-to-text", model=model_name, batch_size=max_batch_size)
def chat_completion(
self,
messages,
temperature=0.6,
top_p: float = 0.9,
max_gen_len=None,
):
generation_output = self.pipe(
text=messages, temperature=temperature, top_p=top_p, max_new_tokens=max_gen_len, return_full_text=False
)
if len(generation_output) == 1:
return CompletionMessage(content=generation_output[0]["generated_text"])
return [CompletionMessage(content=output[0]["generated_text"]) for output in generation_output]
class GeneratorText:
def __init__(self, model_name):
from transformers import AutoModelForImageTextToText, AutoProcessor
self.processor = AutoProcessor.from_pretrained(model_name)
self.model = AutoModelForImageTextToText.from_pretrained(model_name)
def text_completion(
self,
content: Union[str, Image.Image, Dict, List[Dict]],
temperature: float = 0.6,
top_p: float = 0.9,
max_gen_len=None,
):
images = []
text = encode_content(content, images, self.processor.image_token)
model_input = self.processor(text=text, images=images or None, return_tensors="pt", add_special_tokens=False)
tokens = self.model.generate(**model_input, temperature=temperature, top_p=top_p, max_new_tokens=max_gen_len)[0]
tokens = tokens[model_input["input_ids"].shape[-1] :]
return self.processor.decode(tokens, skip_special_tokens=True)
|