cksghl1004/cpp_moondream2
0200
1import torch2import torch.nn as nn3import random4 5from typing import Literal, Tuple, TypedDict, Union, Dict, Any, Optional, List6from PIL import Image7from dataclasses import dataclass8from tokenizers import Tokenizer9 10from .config import MoondreamConfig11from .image_crops import reconstruct_from_crops12from .vision import vision_encoder, vision_projection, prepare_crops, build_vision_model13from .text import build_text_model, text_encoder, lm_head, text_decoder14from .region import (15 decode_coordinate,16 encode_coordinate,17 decode_size,18 encode_size,19 encode_spatial_refs,20 SpatialRefs,21)22from .layers import QuantizedLinear23from .lora import variant_state_dict24from .utils import remove_outlier_points25 26ImageEncodingSettings = TypedDict(27 "ImageEncodingSettings",28 {"variant": str},29 total=False,30)31 32TextSamplingSettings = TypedDict(33 "TextSamplingSettings",34 {35 "max_tokens": int,36 "temperature": float,37 "top_p": float,38 "variant": str,39 },40 total=False,41)42 43ObjectSamplingSettings = TypedDict(44 "ObjectSamplingSettings",45 {"max_objects": int, "variant": str},46 total=False,47)48 49 50DEFAULT_MAX_TOKENS = 76851DEFAULT_TEMPERATURE = 0.552DEFAULT_TOP_P = 0.353DEFAULT_MAX_OBJECTS = 5054 55 56@dataclass(frozen=True)57class EncodedImage:58 pos: int59 caches: List[Tuple[torch.Tensor, torch.Tensor]]60 61 62class KVCache(nn.Module):63 64 def __init__(self, n_heads, n_kv_heads, max_context, dim, device, dtype):65 super().__init__()66 cache_shape = (1, n_kv_heads, max_context, dim // n_heads)67 self.register_buffer(68 "k_cache", torch.zeros(*cache_shape, device=device, dtype=dtype)69 )70 self.register_buffer(71 "v_cache", torch.zeros(*cache_shape, device=device, dtype=dtype)72 )73 74 def update(self, pos_ids, k, v):75 kout, vout = self.k_cache, self.v_cache76 kout[:, :, pos_ids, :] = k77 vout[:, :, pos_ids, :] = v78 return kout, vout79 80 81class MoondreamModel(nn.Module):82 83 def __init__(84 self, config: MoondreamConfig, dtype=torch.bfloat16, setup_caches=True85 ):86 super().__init__()87 self.config = config88 89 self.tokenizer = Tokenizer.from_pretrained("moondream/starmie-v1")90 self.vision = build_vision_model(config.vision, dtype)91 self.text = build_text_model(config.text, dtype)92 93 # Region Model94 linear_cls = (95 QuantizedLinear if config.region.group_size is not None else nn.Linear96 )97 self.region = nn.ModuleDict(98 {99 "coord_encoder": linear_cls(100 config.region.coord_feat_dim, config.region.dim, dtype=dtype101 ),102 "coord_decoder": nn.ModuleDict(103 {104 "fc1": linear_cls(105 config.region.dim, config.region.inner_dim, dtype=dtype106 ),107 "fc2": linear_cls(108 config.region.inner_dim,109 config.region.coord_out_dim,110 dtype=dtype,111 ),112 }113 ),114 "size_encoder": linear_cls(115 config.region.size_feat_dim, config.region.dim, dtype=dtype116 ),117 "size_decoder": nn.ModuleDict(118 {119 "fc1": linear_cls(120 config.region.dim, config.region.inner_dim, dtype=dtype121 ),122 "fc2": linear_cls(123 config.region.inner_dim,124 config.region.size_out_dim,125 dtype=dtype,126 ),127 }128 ),129 }130 )131 self.region.coord_features = nn.Parameter(132 torch.empty(config.region.coord_feat_dim // 2, 1, dtype=dtype).T133 )134 self.region.size_features = nn.Parameter(135 torch.empty(config.region.size_feat_dim // 2, 2, dtype=dtype).T136 )137 138 attn_mask = torch.tril(139 torch.ones(140 1, 1, config.text.max_context, config.text.max_context, dtype=torch.bool141 )142 )143 patch_w = config.vision.crop_size // config.vision.enc_patch_size144 prefix_attn_len = 1 + patch_w**2145 attn_mask[..., :prefix_attn_len, :prefix_attn_len] = 1146 self.register_buffer("attn_mask", attn_mask, persistent=False)147 148 # Initialize KV caches.149 if setup_caches:150 self._setup_caches()151 152 def _setup_caches(self):153 c = self.config.text154 for b in self.text.blocks:155 b.kv_cache = KVCache(156 c.n_heads,157 c.n_kv_heads,158 c.max_context,159 c.dim,160 device=self.device,161 dtype=self.vision.pos_emb.dtype,162 )163 164 @property165 def device(self):166 return self.vision.pos_emb.device167 168 def _vis_enc(self, x: torch.Tensor):169 return vision_encoder(x, self.vision, self.config.vision)170 171 def _vis_proj(self, g: torch.Tensor, r: torch.Tensor):172 return vision_projection(g, r, self.vision, self.config.vision)173 174 def _prefill(175 self,176 x: torch.Tensor,177 attn_mask: torch.Tensor,178 pos_ids: torch.Tensor,179 lora: Optional[torch.Tensor],180 ):181 return text_decoder(x, self.text, attn_mask, pos_ids, self.config.text, lora)182 183 def _decode_one_tok(184 self,185 x: torch.Tensor,186 attn_mask: torch.Tensor,187 pos_ids: torch.Tensor,188 lora: Optional[torch.Tensor],189 ):190 hidden = text_decoder(x, self.text, attn_mask, pos_ids, self.config.text, lora)191 logits = lm_head(hidden, self.text)192 return logits, hidden193 194 def compile(self):195 for module in self.modules():196 if isinstance(module, QuantizedLinear):197 module.unpack()198 199 # TODO: vision_projection is not being compiled200 self._vis_enc = torch.compile(self._vis_enc, fullgraph=True)201 self._prefill = torch.compile(self._prefill, fullgraph=True)202 self._decode_one_tok = torch.compile(203 self._decode_one_tok, fullgraph=True, mode="reduce-overhead"204 )205 206 def _run_vision_encoder(self, image: Image.Image) -> torch.Tensor:207 all_crops, tiling = prepare_crops(image, self.config.vision, device=self.device)208 209 torch._dynamo.mark_dynamic(all_crops, 0)210 211 outputs = self._vis_enc(all_crops)212 213 global_features = outputs[0]214 local_features = outputs[1:].view(215 -1,216 self.config.vision.enc_n_layers,217 self.config.vision.enc_n_layers,218 self.config.vision.enc_dim,219 )220 221 reconstructed = reconstruct_from_crops(222 local_features,223 tiling,224 patch_size=1,225 overlap_margin=self.config.vision.overlap_margin,226 )227 228 return self._vis_proj(global_features, reconstructed)229 230 def encode_image(231 self,232 image: Union[Image.Image, EncodedImage],233 settings: Optional[ImageEncodingSettings] = None,234 ) -> EncodedImage:235 if isinstance(image, EncodedImage):236 return image237 elif not isinstance(image, Image.Image):238 raise ValueError("image must be a PIL Image or EncodedImage")239 240 lora = (241 variant_state_dict(settings["variant"], device=self.device)242 if settings is not None and settings["variant"] is not None243 else None244 )245 246 # Run through text model in addition to the vision encoder, to minimize247 # re-computation if multiple queries are performed on this image.248 with torch.inference_mode():249 img_emb = self._run_vision_encoder(image)250 bos_emb = text_encoder(251 torch.tensor([[self.config.tokenizer.bos_id]], device=self.device),252 self.text,253 )254 inputs_embeds = torch.cat([bos_emb, img_emb[None]], dim=1)255 mask = self.attn_mask[:, :, 0 : inputs_embeds.size(1), :]256 pos_ids = torch.arange(inputs_embeds.size(1), dtype=torch.long)257 self._prefill(inputs_embeds, mask, pos_ids, lora)258 259 return EncodedImage(260 pos=inputs_embeds.size(1),261 caches=[262 (263 b.kv_cache.k_cache[:, :, : inputs_embeds.size(1), :].clone(),264 b.kv_cache.v_cache[:, :, : inputs_embeds.size(1), :].clone(),265 )266 for b in self.text.blocks267 ],268 )269 270 def _apply_top_p(self, probs: torch.Tensor, top_p: float):271 probs_sort, probs_idx = torch.sort(probs, dim=-1, descending=True)272 probs_sum = torch.cumsum(probs_sort, dim=-1)273 mask = probs_sum - probs_sort > top_p274 probs_sort[mask] = 0.0275 probs_sort.div_(probs_sort.sum(dim=-1, keepdim=True))276 next_probs = torch.zeros_like(probs)277 next_probs.scatter_(dim=-1, index=probs_idx, src=probs_sort)278 return next_probs279 280 def _prefill_prompt(281 self,282 prompt_tokens: torch.Tensor,283 pos: int,284 temperature: float,285 top_p: float,286 spatial_refs: Optional[SpatialRefs] = None,287 attn_mask: Optional[torch.Tensor] = None,288 lora: Optional[dict] = None,289 ):290 with torch.inference_mode():291 prompt_emb = text_encoder(prompt_tokens, self.text)292 293 if spatial_refs:294 encoded_refs = encode_spatial_refs(spatial_refs, self.region)295 prompt_emb[prompt_tokens == self.config.tokenizer.coord_id] = (296 encoded_refs["coords"]297 )298 if encoded_refs["sizes"] is not None:299 prompt_emb[prompt_tokens == self.config.tokenizer.size_id] = (300 encoded_refs["sizes"]301 )302 303 torch._dynamo.mark_dynamic(prompt_emb, 1)304 305 if attn_mask is None:306 attn_mask = self.attn_mask307 308 mask = attn_mask[:, :, pos : pos + prompt_emb.size(1), :]309 pos_ids = torch.arange(pos, pos + prompt_emb.size(1), dtype=torch.long)310 hidden_BC = self._prefill(prompt_emb, mask, pos_ids, lora)311 logits_BV = lm_head(hidden_BC, self.text)312 313 if temperature == 0:314 next_token = torch.argmax(logits_BV, dim=-1).unsqueeze(1)315 else:316 probs = torch.softmax(logits_BV / temperature, dim=-1)317 probs = self._apply_top_p(probs, top_p)318 next_token = torch.multinomial(probs, num_samples=1)319 320 pos = pos + prompt_emb.size(1)321 return logits_BV, hidden_BC, next_token, pos322 323 def _generate_reasoning(324 self,325 prompt_tokens,326 pos,327 settings: Optional[TextSamplingSettings] = None,328 spatial_refs: Optional[SpatialRefs] = None,329 attn_mask: Optional[torch.Tensor] = None,330 ) -> Tuple[int, str, List[dict]]:331 max_tokens = (332 settings.get("max_tokens", DEFAULT_MAX_TOKENS)333 if settings334 else DEFAULT_MAX_TOKENS335 )336 temperature = (337 settings.get("temperature", DEFAULT_TEMPERATURE)338 if settings339 else DEFAULT_TEMPERATURE340 )341 lora = (342 variant_state_dict(settings["variant"], device=self.device)343 if settings is not None and "variant" in settings344 else None345 )346 347 top_p = settings.get("top_p", DEFAULT_TOP_P) if settings else DEFAULT_TOP_P348 eos_id = self.config.tokenizer.answer_id349 350 _, last_hidden_BC, next_token, pos = self._prefill_prompt(351 prompt_tokens,352 pos,353 temperature,354 top_p,355 spatial_refs,356 attn_mask=attn_mask,357 lora=lora,358 )359 360 text_token_chunks = [[]]361 grounding_chunks = [[]]362 363 mask = torch.zeros(1, 1, 2048, device=self.device, dtype=torch.bool)364 mask[:, :, :pos] = 1365 pos_ids = torch.tensor([pos], device=self.device, dtype=torch.long)366 generated_tokens = 0367 368 while (369 next_token_id := next_token.item()370 ) != eos_id and generated_tokens < max_tokens:371 if (372 next_token_id == self.config.tokenizer.start_ground_points_id373 or next_token_id == self.config.tokenizer.end_ground_id374 ):375 text_token_chunks.append([])376 grounding_chunks.append([])377 378 text_token_chunks[-1].append(next_token_id)379 380 with torch.inference_mode():381 if next_token_id == self.config.tokenizer.coord_id:382 coord_logits = decode_coordinate(last_hidden_BC, self.region)383 coord = torch.argmax(coord_logits, dim=-1) / coord_logits.size(-1)384 grounding_chunks[-1].append(coord.item())385 386 next_emb = encode_coordinate(387 coord.to(dtype=coord_logits.dtype), self.region388 ).unsqueeze(0)389 else:390 next_emb = text_encoder(next_token, self.text)391 392 mask[:, :, pos], pos_ids[0] = 1, pos393 394 logits_BV, last_hidden_BC = self._decode_one_tok(395 next_emb, mask, pos_ids, lora396 )397 logits_BV[:, self.config.tokenizer.eos_id] = float("-inf")398 logits_BV[:, self.config.tokenizer.size_id] = float("-inf")399 400 pos += 1401 402 if temperature == 0:403 next_token = torch.argmax(logits_BV, dim=-1).unsqueeze(1) # (1, 1)404 else:405 probs = torch.softmax(logits_BV / temperature, dim=-1) # (1, V)406 probs = self._apply_top_p(probs, top_p)407 next_token = torch.multinomial(probs, num_samples=1) # (1, 1)408 409 generated_tokens += 1410 411 text_chunks = [412 self.tokenizer.decode(chunk_tokens) for chunk_tokens in text_token_chunks413 ]414 text = "".join(text_chunks)415 416 start_idx = 0417 grounding = []418 for text_chunk, grounding_chunk in zip(text_chunks, grounding_chunks):419 if len(grounding_chunk) > 1:420 points = []421 for i in range(0, len(grounding_chunk) - (len(grounding_chunk) % 2), 2):422 points.append((grounding_chunk[i], grounding_chunk[i + 1]))423 grounding.append(424 {425 "start_idx": start_idx,426 "end_idx": start_idx + len(text_chunk),427 "points": points,428 }429 )430 start_idx += len(text_chunk)431 432 return pos, text, grounding433 434 def _generate_answer(435 self,436 prompt_tokens: torch.Tensor,437 pos: int,438 settings: Optional[TextSamplingSettings] = None,439 spatial_refs: Optional[SpatialRefs] = None,440 eos_id: Optional[int] = None,441 attn_mask: Optional[torch.Tensor] = None,442 ):443 max_tokens = (444 settings.get("max_tokens", DEFAULT_MAX_TOKENS)445 if settings446 else DEFAULT_MAX_TOKENS447 )448 temperature = (449 settings.get("temperature", DEFAULT_TEMPERATURE)450 if settings451 else DEFAULT_TEMPERATURE452 )453 top_p = settings.get("top_p", DEFAULT_TOP_P) if settings else DEFAULT_TOP_P454 eos_id = eos_id if eos_id is not None else self.config.tokenizer.eos_id455 lora = (456 variant_state_dict(settings["variant"], device=self.device)457 if settings is not None and "variant" in settings458 else None459 )460 461 _, _, next_token, pos = self._prefill_prompt(462 prompt_tokens,463 pos,464 temperature,465 top_p,466 spatial_refs,467 attn_mask=attn_mask,468 lora=lora,469 )470 471 def generator(next_token, pos):472 mask = torch.zeros(1, 1, 2048, device=self.device, dtype=torch.bool)473 mask[:, :, :pos] = 1474 pos_ids = torch.tensor([pos], device=self.device, dtype=torch.long)475 generated_tokens = 0476 477 # For properly handling token streaming with Unicode478 token_cache = []479 print_len = 0480 481 while (482 next_token_id := next_token.item()483 ) != eos_id and generated_tokens < max_tokens:484 # Add token to our cache485 token_cache.append(next_token_id)486 487 # Decode all tokens collected so far488 text = self.tokenizer.decode(token_cache)489 490 # After a newline, we flush the cache completely491 if text.endswith("\n"):492 printable_text = text[print_len:]493 token_cache = []494 print_len = 0495 if printable_text:496 yield printable_text497 # If the last token is a CJK character, we can safely print it498 elif len(text) > 0 and _is_cjk_char(ord(text[-1])):499 printable_text = text[print_len:]500 print_len += len(printable_text)501 if printable_text:502 yield printable_text503 # Otherwise, only yield up to the last space to avoid cutting words504 else:505 last_space_idx = text.rfind(" ", print_len)506 if last_space_idx >= print_len:507 printable_text = text[print_len : last_space_idx + 1]508 print_len += len(printable_text)509 if printable_text:510 yield printable_text511 512 with torch.inference_mode():513 next_emb = text_encoder(next_token, self.text)514 mask[:, :, pos], pos_ids[0] = 1, pos515 516 logits_BV, _ = self._decode_one_tok(next_emb, mask, pos_ids, lora)517 logits_BV[:, self.config.tokenizer.answer_id] = float("-inf")518 519 pos += 1520 521 if temperature == 0:522 next_token = torch.argmax(logits_BV, dim=-1).unsqueeze(523 1524 ) # (1, 1)525 else:526 probs = torch.softmax(logits_BV / temperature, dim=-1) # (1, V)527 probs = self._apply_top_p(probs, top_p)528 next_token = torch.multinomial(probs, num_samples=1) # (1, 1)529 530 generated_tokens += 1531 532 # Flush any remaining text in the cache533 if token_cache:534 text = self.tokenizer.decode(token_cache)535 printable_text = text[print_len:]536 if printable_text:537 yield printable_text538 539 return generator(next_token, pos)540 541 def query(542 self,543 image: Optional[Union[Image.Image, EncodedImage]] = None,544 question: str = None,545 reasoning: bool = False,546 spatial_refs: Optional[SpatialRefs] = None,547 stream: bool = False,548 settings: Optional[TextSamplingSettings] = None,549 ):550 if self.config.tokenizer.templates["query"] is None:551 raise NotImplementedError("Model does not support querying.")552 553 if question is None:554 raise ValueError("question must be provided.")555 556 if spatial_refs and image is None:557 raise ValueError("spatial_refs can only be used with an image.")558 559 attn_mask = self.attn_mask560 if image is not None:561 image = self.encode_image(image, settings)562 self.load_encoded_image(image)563 pos = image.pos564 prompt_toks = self.config.tokenizer.templates["query"]["prefix"]565 else:566 self._setup_caches()567 pos = 0568 prompt_toks = [569 self.config.tokenizer.bos_id570 ] + self.config.tokenizer.templates["query"]["prefix"]571 max_context = self.config.text.max_context572 attn_mask = torch.tril(573 torch.ones(1, 1, max_context, max_context, dtype=torch.bool)574 ).to(self.device)575 576 spatial_toks = []577 if spatial_refs:578 for ref in spatial_refs:579 coord_id = self.config.tokenizer.coord_id580 size_id = self.config.tokenizer.size_id581 if len(ref) == 2:582 spatial_toks.extend([coord_id, coord_id])583 else:584 spatial_toks.extend([coord_id, coord_id, size_id])585 586 prompt_tokens = [587 prompt_toks588 + spatial_toks589 + self.tokenizer.encode(question).ids590 + self.config.tokenizer.templates["query"]["suffix"]591 ]592 593 if reasoning:594 prompt_tokens[0] += [self.config.tokenizer.thinking_id]595 prompt_tokens = torch.tensor(prompt_tokens, device=self.device)596 pos, reasoning_text, reasoning_grounding = self._generate_reasoning(597 prompt_tokens, pos, settings, spatial_refs, attn_mask=attn_mask598 )599 prompt_tokens = [self.config.tokenizer.templates["query"]["suffix"]]600 reasoning_dict = {601 "reasoning": {"text": reasoning_text, "grounding": reasoning_grounding}602 }603 else:604 prompt_tokens[0] += self.config.tokenizer.templates["query"]["suffix"]605 reasoning_dict = {}606 607 prompt_tokens = torch.tensor(prompt_tokens, device=self.device)608 609 def generator():610 for token in self._generate_answer(611 prompt_tokens, pos, settings, spatial_refs, attn_mask=attn_mask612 ):613 yield token614 615 if stream:616 return {**reasoning_dict, "answer": generator()}617 else:618 return {**reasoning_dict, "answer": "".join(list(generator()))}619 620 def load_encoded_image(self, encoded_image: EncodedImage):621 for b, (k, v) in zip(self.text.blocks, encoded_image.caches):622 b.kv_cache.k_cache[:, :, : k.size(2), :] = k623 b.kv_cache.v_cache[:, :, : v.size(2), :] = v624 625 def caption(626 self,627 image: Union[Image.Image, EncodedImage],628 length: Literal["normal", "short", "long"] = "normal",629 stream: bool = False,630 settings: Optional[TextSamplingSettings] = None,631 ):632 if self.config.tokenizer.templates["caption"] is None:633 raise NotImplementedError("Model does not support captioning.")634 if length not in self.config.tokenizer.templates["caption"]:635 raise ValueError(f"Model does not support caption length '{length}'.")636 637 image = self.encode_image(image, settings)638 self.load_encoded_image(image)639 640 prompt_tokens = torch.tensor(641 [self.config.tokenizer.templates["caption"][length]], device=self.device642 )643 644 def generator():645 for token in self._generate_answer(prompt_tokens, image.pos, settings):646 yield token647 648 if stream:649 return {"caption": generator()}650 else:651 return {"caption": "".join(list(generator()))}652 653 def _generate_points(654 self,655 hidden: torch.Tensor,656 next_token: torch.Tensor,657 pos: int,658 include_size: bool = True,659 max_objects: int = DEFAULT_MAX_OBJECTS,660 lora: Optional[dict] = None,661 ):662 out = []663 mask = torch.zeros(1, 1, 2048, device=self.device, dtype=torch.bool)664 mask[:, :, :pos] = 1665 pos_ids = torch.tensor([pos], device=self.device, dtype=torch.long)666 667 with torch.inference_mode():668 while (669 next_token.item() != self.config.tokenizer.eos_id670 and len(out) < max_objects671 ):672 x_logits = decode_coordinate(hidden, self.region)673 x_center = torch.argmax(x_logits, dim=-1) / x_logits.size(-1)674 next_emb = encode_coordinate(675 x_center.to(dtype=x_logits.dtype), self.region676 ).unsqueeze(0)677 678 # Decode y-coordinate679 mask[:, :, pos], pos_ids[0] = 1, pos680 _, hidden = self._decode_one_tok(next_emb, mask, pos_ids, lora)681 pos += 1682 y_logits = decode_coordinate(hidden, self.region)683 y_center = torch.argmax(y_logits, dim=-1) / y_logits.size(-1)684 next_emb = encode_coordinate(685 y_center.to(dtype=y_logits.dtype), self.region686 ).unsqueeze(0)687 688 # Decode size689 if include_size:690 mask[:, :, pos], pos_ids[0] = 1, pos691 logits, hidden = self._decode_one_tok(next_emb, mask, pos_ids, lora)692 pos += 1693 size_logits = decode_size(hidden, self.region)694 695 # Get bin indices from the logits696 w_bin = torch.argmax(size_logits[0], dim=-1)697 h_bin = torch.argmax(size_logits[1], dim=-1)698 699 # Convert from bin indices to actual size values using the inverse of the log-scale mapping700 # Formula: size = 2^((bin / 1023.0) * 10.0 - 10.0)701 w = torch.pow(2.0, (w_bin.float() / 1023.0) * 10.0 - 10.0)702 h = torch.pow(2.0, (h_bin.float() / 1023.0) * 10.0 - 10.0)703 704 next_emb = (705 encode_size(706 torch.tensor(707 [w, h], device=self.device, dtype=size_logits.dtype708 ),709 self.region,710 )711 .unsqueeze(0)712 .unsqueeze(0)713 )714 715 # Add object716 out.append(717 {718 "x_min": x_center.item() - w.item() / 2,719 "y_min": y_center.item() - h.item() / 2,720 "x_max": x_center.item() + w.item() / 2,721 "y_max": y_center.item() + h.item() / 2,722 }723 )724 else:725 out.append({"x": x_center.item(), "y": y_center.item()})726 727 # Decode next token (x-coordinate, or eos)728 mask[:, :, pos], pos_ids[0] = 1, pos729 logits, hidden = self._decode_one_tok(next_emb, mask, pos_ids, lora)730 pos += 1731 next_token = torch.argmax(logits, dim=-1)732 733 return out734 735 def detect(736 self,737 image: Union[Image.Image, EncodedImage],738 object: str,739 settings: Optional[ObjectSamplingSettings] = None,740 ):741 if self.config.tokenizer.templates["detect"] is None:742 raise NotImplementedError("Model does not support object detection.")743 744 image = self.encode_image(image, settings)745 self.load_encoded_image(image)746 747 prompt_tokens = torch.tensor(748 [749 self.config.tokenizer.templates["detect"]["prefix"]750 + self.tokenizer.encode(" " + object).ids751 + self.config.tokenizer.templates["detect"]["suffix"]752 ],753 device=self.device,754 )755 756 lora = (757 variant_state_dict(settings["variant"], device=self.device)758 if settings is not None and "variant" in settings759 else None760 )761 762 _, hidden, next_token, pos = self._prefill_prompt(763 prompt_tokens, image.pos, temperature=0, top_p=0, lora=lora764 )765 hidden = hidden[:, -1:, :]766 767 max_objects = (768 settings.get("max_objects", DEFAULT_MAX_OBJECTS)769 if settings770 else DEFAULT_MAX_OBJECTS771 )772 objects = self._generate_points(773 hidden,774 next_token,775 pos,776 include_size=True,777 max_objects=max_objects,778 lora=lora,779 )780 781 return {"objects": objects}782 783 def point(784 self,785 image: Union[Image.Image, EncodedImage],786 object: str,787 settings: Optional[ObjectSamplingSettings] = None,788 ):789 if self.config.tokenizer.templates["point"] is None:790 raise NotImplementedError("Model does not support pointing.")791 792 image = self.encode_image(image, settings)793 self.load_encoded_image(image)794 795 prompt_tokens = torch.tensor(796 [797 self.config.tokenizer.templates["point"]["prefix"]798 + self.tokenizer.encode(" " + object).ids799 + self.config.tokenizer.templates["point"]["suffix"]800 ],801 device=self.device,802 )803 804 lora = (805 variant_state_dict(settings["variant"], device=self.device)806 if settings is not None and "variant" in settings807 else None808 )809 810 _, hidden, next_token, pos = self._prefill_prompt(811 prompt_tokens, image.pos, temperature=0, top_p=0, lora=lora812 )813 hidden = hidden[:, -1:, :]814 815 max_objects = (816 settings.get("max_objects", DEFAULT_MAX_OBJECTS)817 if settings818 else DEFAULT_MAX_OBJECTS819 )820 objects = self._generate_points(821 hidden,822 next_token,823 pos,824 include_size=False,825 max_objects=max_objects,826 lora=lora,827 )828 829 return {"points": objects}830 831 def _detect_gaze(832 self,833 image: EncodedImage,834 source: Tuple[float, float],835 force_detect: bool = False,836 ):837 with torch.inference_mode():838 before_emb = text_encoder(839 torch.tensor(840 [self.tokenizer.encode("\n\nPoint:").ids], device=self.device841 ),842 self.text,843 )844 after_emb = text_encoder(845 torch.tensor(846 [self.tokenizer.encode(" gaze\n\n").ids], device=self.device847 ),848 self.text,849 )850 x_emb = encode_coordinate(851 torch.tensor([[[source[0]]]], device=self.device, dtype=torch.bfloat16),852 self.region,853 )854 y_emb = encode_coordinate(855 torch.tensor([[[source[1]]]], device=self.device, dtype=torch.bfloat16),856 self.region,857 )858 859 prompt_emb = torch.cat([before_emb, x_emb, y_emb, after_emb], dim=1)860 861 self.load_encoded_image(image)862 863 mask = self.attn_mask[:, :, image.pos : image.pos + prompt_emb.size(1), :]864 pos_ids = torch.arange(865 image.pos, image.pos + prompt_emb.size(1), dtype=torch.long866 )867 hidden = self._prefill(prompt_emb, mask, pos_ids, lora=None)868 logits = lm_head(hidden, self.text)869 next_token = torch.argmax(logits, dim=-1)870 pos = image.pos + prompt_emb.size(1)871 hidden = hidden[:, -1:, :]872 873 if force_detect:874 next_token = torch.tensor([[0]], device=self.device)875 876 if next_token.item() == self.config.tokenizer.eos_id:877 return None878 879 gaze = self._generate_points(880 hidden, next_token, pos, include_size=False, max_objects=1881 )882 return gaze[0]883 884 def detect_gaze(885 self,886 image: Union[Image.Image, EncodedImage],887 eye: Optional[Tuple[float, float]] = None,888 face: Optional[Dict[str, float]] = None,889 unstable_settings: Dict[str, Any] = {},890 ):891 if "force_detect" in unstable_settings:892 force_detect = unstable_settings["force_detect"]893 else:894 force_detect = False895 896 if "prioritize_accuracy" in unstable_settings:897 prioritize_accuracy = unstable_settings["prioritize_accuracy"]898 else:899 prioritize_accuracy = False900 901 if not prioritize_accuracy:902 if eye is None:903 raise ValueError("eye must be provided when prioritize_accuracy=False")904 image = self.encode_image(image)905 return {"gaze": self._detect_gaze(image, eye, force_detect=force_detect)}906 else:907 if (908 not isinstance(image, Image.Image)909 and "flip_enc_img" not in unstable_settings910 ):911 raise ValueError(912 "image must be a PIL Image when prioritize_accuracy=True, "913 "or flip_enc_img must be provided"914 )915 if face is None:916 raise ValueError("face must be provided when prioritize_accuracy=True")917 918 encoded_image = self.encode_image(image)919 if (920 isinstance(image, Image.Image)921 and "flip_enc_img" not in unstable_settings922 ):923 flipped_pil = image.copy()924 flipped_pil = flipped_pil.transpose(method=Image.FLIP_LEFT_RIGHT)925 encoded_flipped_image = self.encode_image(flipped_pil)926 else:927 encoded_flipped_image = unstable_settings["flip_enc_img"]928 929 N = 10930 931 detections = [932 self._detect_gaze(933 encoded_image,934 (935 random.uniform(face["x_min"], face["x_max"]),936 random.uniform(face["y_min"], face["y_max"]),937 ),938 force_detect=force_detect,939 )940 for _ in range(N)941 ]942 detections = [943 (gaze["x"], gaze["y"]) for gaze in detections if gaze is not None944 ]945 flipped_detections = [946 self._detect_gaze(947 encoded_flipped_image,948 (949 1 - random.uniform(face["x_min"], face["x_max"]),950 random.uniform(face["y_min"], face["y_max"]),951 ),952 force_detect=force_detect,953 )954 for _ in range(N)955 ]956 detections.extend(957 [958 (1 - gaze["x"], gaze["y"])959 for gaze in flipped_detections960 if gaze is not None961 ]962 )963 964 if len(detections) < N:965 return {"gaze": None}966 967 detections = remove_outlier_points(detections)968 mean_gaze = (969 sum(gaze[0] for gaze in detections) / len(detections),970 sum(gaze[1] for gaze in detections) / len(detections),971 )972 973 return {"gaze": {"x": mean_gaze[0], "y": mean_gaze[1]}}974 975 976def _is_cjk_char(cp):977 """Checks whether CP is the codepoint of a CJK character."""978 # This defines a "chinese character" as anything in the CJK Unicode block:979 # https://en.wikipedia.org/wiki/CJK_Unified_Ideographs_(Unicode_block)980 if (981 (cp >= 0x4E00 and cp <= 0x9FFF)982 or (cp >= 0x3400 and cp <= 0x4DBF)983 or (cp >= 0x2F800 and cp <= 0x2FA1F)984 ):985 return True986 return False987 