Team Ai
Modelpublic

cksghl1004/cpp_moondream2

sourceHugging Faceapache-2.0updated 1y agoView on Hugging Face
0likes200downloads
moondream.py987 linesDownload Raw Back to root
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