Instructions to use ruotian/SelectGround-8B with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- PEFT
How to use ruotian/SelectGround-8B with PEFT:
from peft import PeftModel from transformers import AutoModelForCausalLM base_model = AutoModelForCausalLM.from_pretrained("Qwen/Qwen3-VL-8B-Instruct") model = PeftModel.from_pretrained(base_model, "ruotian/SelectGround-8B") - Notebooks
- Google Colab
- Kaggle
| from __future__ import annotations | |
| from collections import defaultdict | |
| from dataclasses import dataclass | |
| import math | |
| from pathlib import Path | |
| from typing import Any | |
| from PIL import Image | |
| import torch | |
| from transformers.cache_utils import DynamicCache | |
| from selectground import SelectGround, _map_crop, _prediction | |
| GRID_CENTERS = ((0.3, 0.3), (0.7, 0.3), (0.3, 0.7), (0.7, 0.7)) | |
| class _Prefix: | |
| cache: DynamicCache | |
| logits: torch.Tensor | |
| position_ids: torch.Tensor | |
| attention_mask: torch.Tensor | |
| length: int | |
| class SelfContrastGrounder: | |
| """Training-free self-contrastive grounding with six visual prefills.""" | |
| def __init__(self, checkpoint: str = "ruotian/SelectGround-8B") -> None: | |
| self.grounder = SelectGround(checkpoint) | |
| def predict( | |
| self, | |
| image: str | Path | Image.Image, | |
| instruction: str, | |
| *, | |
| variant: str = "full", | |
| ) -> dict[str, Any]: | |
| variants = { | |
| "full": (GRID_CENTERS, True, True, 1.0), | |
| "no_latent_distractors": ((), True, True, 1.0), | |
| "one_latent_distractor": (GRID_CENTERS[:1], True, True, 1.0), | |
| "no_recurrent_anchor": (GRID_CENTERS, False, True, 1.0), | |
| "no_cross_view_evidence": (GRID_CENTERS, True, False, 1.0), | |
| "no_anchor_proximity": (GRID_CENTERS, True, True, 0.0), | |
| } | |
| if variant not in variants: | |
| raise ValueError(f"unknown self-contrast variant: {variant}") | |
| grid_centers, use_anchor, use_evidence, proximity_weight = variants[variant] | |
| source = ( | |
| Image.open(image).convert("RGB") | |
| if not isinstance(image, Image.Image) | |
| else image.convert("RGB") | |
| ) | |
| full_box = (0, 0, source.width, source.height) | |
| views: dict[str, tuple[tuple[int, int, int, int], Image.Image]] = { | |
| "full": (full_box, source) | |
| } | |
| candidates = [] | |
| prefixes = {} | |
| p0, prefixes["full"] = self._observe(source, instruction) | |
| candidates.append(self._candidate("p0", p0, full_box, source.size)) | |
| if use_anchor and p0["point"] is not None: | |
| box = _crop_box(tuple(p0["point"]), source.size, 0.40) | |
| views["q0"] = (box, _view(source, box)) | |
| for index, center in enumerate(grid_centers): | |
| point = (center[0] * source.width, center[1] * source.height) | |
| box = _crop_box(point, source.size, 0.60) | |
| views[f"grid_{index}"] = (box, _view(source, box)) | |
| for name, (box, view) in tuple(views.items())[1:]: | |
| prediction, prefixes[name] = self._observe(view, instruction) | |
| candidates.append(self._candidate(name, prediction, box, source.size)) | |
| evidence = {} | |
| for view_name, (box, _) in views.items(): | |
| visible = { | |
| candidate["name"]: response | |
| for candidate in candidates | |
| if candidate["point"] is not None | |
| and (response := _response(candidate["point"], box)) is not None | |
| } | |
| scores = self._score(prefixes.pop(view_name), list(visible.values())) | |
| evidence[view_name] = { | |
| name: {"response": response, **scores[response]} | |
| for name, response in visible.items() | |
| } | |
| selected = _select( | |
| candidates, | |
| evidence, | |
| source.size, | |
| proximity_weight=proximity_weight, | |
| use_evidence=use_evidence, | |
| ) | |
| point = selected["point"] | |
| normalized = ( | |
| [1000 * point[0] / source.width, 1000 * point[1] / source.height] | |
| if point is not None | |
| else None | |
| ) | |
| return { | |
| "method": "SelectGround+SelfContrast", | |
| "variant": variant, | |
| "point": point, | |
| "normalized_point": normalized, | |
| "raw_response": selected["raw_response"], | |
| "selected_candidate": selected["name"], | |
| } | |
| def _candidate( | |
| self, | |
| name: str, | |
| prediction: dict[str, Any], | |
| box: tuple[int, int, int, int], | |
| source_size: tuple[int, int], | |
| ) -> dict[str, Any]: | |
| mapped = ( | |
| prediction | |
| if name == "p0" or prediction["point"] is None | |
| else _map_crop(prediction, box, source_size, 2.0) | |
| ) | |
| return { | |
| "name": name, | |
| "point": mapped["point"], | |
| "source_view": "full" if name == "p0" else name, | |
| "raw_response": prediction["raw_response"], | |
| } | |
| def _observe( | |
| self, image: Image.Image, instruction: str | |
| ) -> tuple[dict[str, Any], _Prefix]: | |
| inputs = self.grounder._inputs(image, instruction, False) | |
| input_ids = inputs["input_ids"] | |
| length = int(input_ids.shape[1]) | |
| position_ids, _ = self.grounder.core.get_rope_index( | |
| input_ids, | |
| inputs.get("image_grid_thw"), | |
| inputs.get("video_grid_thw"), | |
| attention_mask=inputs.get("attention_mask"), | |
| ) | |
| cache = DynamicCache(config=self.grounder.core.language_model.config) | |
| output = self.grounder.model( | |
| **inputs, | |
| past_key_values=cache, | |
| position_ids=position_ids, | |
| cache_position=torch.arange(length, device=self.grounder.device), | |
| use_cache=True, | |
| logits_to_keep=1, | |
| ) | |
| logits = output.logits[:, -1, :].detach() | |
| raw = self.grounder._decode( | |
| logits, cache, position_ids[:, :, -1:] + 1, None | |
| ) | |
| cache.crop(length) | |
| if cache.get_seq_length() != length: | |
| raise RuntimeError("could not restore the visual prefix after decoding") | |
| return ( | |
| _prediction(raw, image.size, integer=False), | |
| _Prefix(cache, logits, position_ids, inputs["attention_mask"], length), | |
| ) | |
| def _score( | |
| self, prefix: _Prefix, responses: list[str] | |
| ) -> dict[str, dict[str, float | int]]: | |
| unique = list(dict.fromkeys(responses)) | |
| if not unique: | |
| return {} | |
| encoded = [_token_ids(self.grounder, response) for response in unique] | |
| first = torch.log_softmax(prefix.logits.float(), -1) | |
| logps = [[float(first[0, values[0]])] for values in encoded] | |
| maximum = max(map(len, encoded)) | |
| if maximum > 1: | |
| tokenizer = self.grounder.processor.tokenizer | |
| pad = tokenizer.pad_token_id or tokenizer.eos_token_id | |
| continuation = torch.full( | |
| (len(encoded), maximum - 1), | |
| int(pad), | |
| dtype=torch.long, | |
| device=self.grounder.device, | |
| ) | |
| mask = torch.zeros_like(continuation, dtype=torch.bool) | |
| for index, values in enumerate(encoded): | |
| if len(values) > 1: | |
| continuation[index, : len(values) - 1] = torch.tensor( | |
| values[:-1], device=self.grounder.device | |
| ) | |
| mask[index, : len(values) - 1] = True | |
| prefix.cache.batch_repeat_interleave(len(encoded)) | |
| positions = prefix.position_ids.repeat_interleave(len(encoded), dim=-2) | |
| offsets = torch.arange(maximum - 1, device=self.grounder.device).view( | |
| *([1] * (positions.ndim - 1)), -1 | |
| ) | |
| output = self.grounder.model( | |
| input_ids=continuation, | |
| past_key_values=prefix.cache, | |
| attention_mask=torch.cat( | |
| (prefix.attention_mask.repeat(len(encoded), 1), mask.long()), 1 | |
| ), | |
| position_ids=positions[..., -1:] + 1 + offsets, | |
| cache_position=torch.arange( | |
| prefix.length, | |
| prefix.length + maximum - 1, | |
| device=self.grounder.device, | |
| ), | |
| use_cache=True, | |
| ) | |
| for index, values in enumerate(encoded): | |
| if len(values) <= 1: | |
| continue | |
| logits = output.logits[index, : len(values) - 1].float() | |
| labels = torch.tensor(values[1:], device=self.grounder.device) | |
| selected = torch.log_softmax(logits, -1).gather(1, labels[:, None])[:, 0] | |
| logps[index].extend(float(value) for value in selected) | |
| return { | |
| response: { | |
| "token_count": len(values), | |
| "mean_logprob": sum(values_logps) / len(values), | |
| } | |
| for response, values, values_logps in zip(unique, encoded, logps, strict=True) | |
| } | |
| def _crop_box( | |
| point: tuple[float, float], size: tuple[int, int], fraction: float | |
| ) -> tuple[int, int, int, int]: | |
| width, height = size | |
| crop_width = min(width, max(320, round(fraction * width))) | |
| crop_height = min(height, max(320, round(fraction * height))) | |
| left = round(min(max(0.0, point[0] - crop_width / 2), width - crop_width)) | |
| top = round(min(max(0.0, point[1] - crop_height / 2), height - crop_height)) | |
| return left, top, left + crop_width, top + crop_height | |
| def _view(source: Image.Image, box: tuple[int, int, int, int]) -> Image.Image: | |
| crop = source.crop(box) | |
| return crop.resize( | |
| (2 * crop.width, 2 * crop.height), Image.Resampling.LANCZOS | |
| ) | |
| def _response(point: list[float], box: tuple[int, int, int, int]) -> str | None: | |
| left, top, right, bottom = box | |
| x, y = map(float, point) | |
| if not left <= x < right or not top <= y < bottom: | |
| return None | |
| return ( | |
| f"[{round(1000 * (x - left) / (right - left))}," | |
| f"{round(1000 * (y - top) / (bottom - top))}]" | |
| ) | |
| def _token_ids(grounder: SelectGround, response: str) -> list[int]: | |
| values = grounder.processor.tokenizer( | |
| response, add_special_tokens=False | |
| )["input_ids"] | |
| if values and isinstance(values[0], list): | |
| values = values[0] | |
| result = [int(value) for value in values] | |
| if not result: | |
| raise ValueError(f"empty tokenization for {response!r}") | |
| return result | |
| def _zscore(values: list[float]) -> list[float]: | |
| mean = sum(values) / len(values) | |
| std = math.sqrt(sum((value - mean) ** 2 for value in values) / len(values)) | |
| return [(value - mean) / max(std, 1e-6) for value in values] | |
| def _select( | |
| candidates: list[dict[str, Any]], | |
| evidence: dict[str, dict[str, dict[str, Any]]], | |
| size: tuple[int, int], | |
| *, | |
| proximity_weight: float, | |
| use_evidence: bool, | |
| ) -> dict[str, Any]: | |
| eligible = [candidate for candidate in candidates if candidate["point"] is not None] | |
| if not eligible: | |
| return candidates[0] | |
| by_name = {candidate["name"]: candidate for candidate in eligible} | |
| accumulated = defaultdict(list) | |
| for view_name, values in evidence.items(): | |
| unique = {} | |
| for name, value in values.items(): | |
| if name in by_name: | |
| unique.setdefault(value["response"], float(value["mean_logprob"])) | |
| if not unique: | |
| continue | |
| normalized = dict(zip(unique, _zscore(list(unique.values())), strict=True)) | |
| for name, value in values.items(): | |
| if name in by_name and by_name[name]["source_view"] != view_name: | |
| accumulated[name].append(normalized[value["response"]]) | |
| if use_evidence: | |
| eligible = [candidate for candidate in eligible if accumulated[candidate["name"]]] | |
| if not eligible: | |
| return candidates[0] | |
| likelihood = _zscore( | |
| [ | |
| sum(accumulated[candidate["name"]]) | |
| / len(accumulated[candidate["name"]]) | |
| for candidate in eligible | |
| ] | |
| ) | |
| else: | |
| likelihood = [0.0] * len(eligible) | |
| by_name = {candidate["name"]: candidate for candidate in eligible} | |
| anchor = by_name.get("q0", by_name.get("p0", eligible[0]))["point"] | |
| width, height = size | |
| proximity = _zscore( | |
| [ | |
| -math.hypot( | |
| (candidate["point"][0] - anchor[0]) / width, | |
| (candidate["point"][1] - anchor[1]) / height, | |
| ) | |
| for candidate in eligible | |
| ] | |
| ) | |
| scores = [ | |
| likelihood[index] + proximity_weight * proximity[index] | |
| for index in range(len(eligible)) | |
| ] | |
| return eligible[max(range(len(eligible)), key=lambda index: (scores[index], -index))] | |