Sentence Similarity
Safetensors
sentence-transformers
Korean
PyLate
modernbert
ColBERT
feature-extraction
Generated from Trainer
loss:MatryoshkaColBERTLoss
korean
matryoshka
retrieval
custom_code
text-embeddings-inference
Instructions to use dragonkue/colbert-ko-0.1b with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- sentence-transformers
How to use dragonkue/colbert-ko-0.1b with sentence-transformers:
from pylate import models queries = [ "Which planet is known as the Red Planet?", "What is the largest planet in our solar system?", ] documents = [ ["Mars is the Red Planet.", "Venus is Earth's twin."], ["Jupiter is the largest planet.", "Saturn has rings."], ] model = models.ColBERT(model_name_or_path="dragonkue/colbert-ko-0.1b") queries_emb = model.encode(queries, is_query=True) docs_emb = model.encode(documents, is_query=False) - Notebooks
- Google Colab
- Kaggle
| """ColBERT-Ko: transformers-only inference for dragonkue/colbert-ko-0.1b. | |
| Usage: | |
| from transformers import AutoModel, AutoTokenizer | |
| tokenizer = AutoTokenizer.from_pretrained("dragonkue/colbert-ko-0.1b") | |
| model = AutoModel.from_pretrained( | |
| "dragonkue/colbert-ko-0.1b", | |
| trust_remote_code=True, | |
| ) | |
| q_embs = model.encode(tokenizer, ["검색 쿼리"], is_query=True) | |
| d_embs = model.encode(tokenizer, ["문서 내용"], is_query=False) | |
| score = model.maxsim(q_embs[0], d_embs[0]) | |
| The model is a ModernBert encoder followed by a 768→128 linear projection | |
| (no bias) producing L2-normalized token-level embeddings. Late-interaction | |
| MaxSim is the similarity function. Query expansion (pad to 32 tokens with | |
| MASK) and document skiplist (punctuation removal) are applied as in the | |
| original ColBERT recipe. | |
| """ | |
| from __future__ import annotations | |
| import os | |
| from typing import Any | |
| import torch | |
| import torch.nn.functional as F | |
| from torch import nn | |
| from transformers import PreTrainedModel | |
| from transformers.modeling_outputs import BaseModelOutput | |
| from transformers.models.modernbert.configuration_modernbert import ModernBertConfig | |
| from transformers.models.modernbert.modeling_modernbert import ModernBertModel | |
| DEFAULT_SKIPLIST_WORDS = list("!\"#$%&'()*+,-./:;<=>?@[\\]^_`{|}~") | |
| class ColBERTKoConfig(ModernBertConfig): | |
| """ModernBert config + ColBERT-specific fields.""" | |
| model_type = "modernbert" # keep ModernBert mapping for the encoder | |
| def __init__( | |
| self, | |
| embedding_dim: int = 128, | |
| query_prefix: str = "[Q] ", | |
| document_prefix: str = "[D] ", | |
| query_length: int = 32, | |
| document_length: int = 2048, | |
| attend_to_expansion_tokens: bool = False, | |
| do_query_expansion: bool = True, | |
| skiplist_words: list[str] | None = None, | |
| **kwargs: Any, | |
| ) -> None: | |
| super().__init__(**kwargs) | |
| self.embedding_dim = embedding_dim | |
| self.query_prefix = query_prefix | |
| self.document_prefix = document_prefix | |
| self.query_length = query_length | |
| self.document_length = document_length | |
| self.attend_to_expansion_tokens = attend_to_expansion_tokens | |
| self.do_query_expansion = do_query_expansion | |
| self.skiplist_words = ( | |
| skiplist_words if skiplist_words is not None else DEFAULT_SKIPLIST_WORDS | |
| ) | |
| class ColBERTKo(PreTrainedModel): | |
| """ModernBert encoder + 768→128 Linear (no bias) + L2 norm.""" | |
| config_class = ColBERTKoConfig | |
| base_model_prefix = "encoder" | |
| _no_split_modules = ["ModernBertEncoderLayer"] | |
| # Linear projection weights are loaded separately from matryoshka_heads.pt | |
| # (or 1_Dense/model.safetensors as a fallback); suppress the spurious | |
| # "missing key" warning from transformers' standard loader. | |
| _keys_to_ignore_on_load_missing = ["linear.weight"] | |
| def __init__(self, config: ColBERTKoConfig) -> None: | |
| super().__init__(config) | |
| self.encoder = ModernBertModel(config) | |
| self.linear = nn.Linear( | |
| config.hidden_size, config.embedding_dim, bias=False | |
| ) | |
| # Matryoshka heads (32/64/96/128) populated lazily from matryoshka_heads.pt. | |
| # Active dim defaults to embedding_dim (128). set_active_dim() switches | |
| # which weight is loaded into self.linear. | |
| self._matryoshka_heads: dict[int, torch.Tensor] = {} | |
| self._active_dim: int = config.embedding_dim | |
| # token id caches resolved at first encode() call | |
| self._query_prefix_id: int | None = None | |
| self._document_prefix_id: int | None = None | |
| self._skiplist_ids: list[int] = [] | |
| # Required by transformers >= 5.x for tied-weight bookkeeping. | |
| self.post_init() | |
| # ---- weight loading (linear comes from matryoshka_heads.pt) ---- | |
| def from_pretrained( # type: ignore[override] | |
| cls, pretrained_model_name_or_path: str | os.PathLike, *args: Any, **kwargs: Any, | |
| ) -> "ColBERTKo": | |
| model: "ColBERTKo" = super().from_pretrained( # type: ignore[assignment] | |
| pretrained_model_name_or_path, *args, **kwargs | |
| ) | |
| model._load_matryoshka_heads(pretrained_model_name_or_path, **kwargs) | |
| model.set_active_dim(model.config.embedding_dim) | |
| return model | |
| def _load_matryoshka_heads( | |
| self, path: str | os.PathLike, **kwargs: Any | |
| ) -> None: | |
| """Load all Matryoshka projection heads (32/64/96/128) from matryoshka_heads.pt. | |
| Falls back to 1_Dense/model.safetensors (128-dim only) if matryoshka_heads.pt | |
| is unavailable. | |
| """ | |
| local_mat = os.path.join(str(path), "matryoshka_heads.pt") | |
| local_1d = os.path.join(str(path), "1_Dense", "model.safetensors") | |
| heads_path: str | None = None | |
| dense_path: str | None = None | |
| if os.path.isdir(str(path)): | |
| if os.path.exists(local_mat): | |
| heads_path = local_mat | |
| if os.path.exists(local_1d): | |
| dense_path = local_1d | |
| else: | |
| from huggingface_hub import hf_hub_download | |
| hf_kwargs: dict[str, Any] = {} | |
| for k in ("revision", "cache_dir", "token", "subfolder"): | |
| if k in kwargs: | |
| hf_kwargs[k] = kwargs[k] | |
| try: | |
| heads_path = hf_hub_download( | |
| repo_id=str(path), filename="matryoshka_heads.pt", **hf_kwargs | |
| ) | |
| except Exception: | |
| heads_path = None | |
| try: | |
| dense_path = hf_hub_download( | |
| repo_id=str(path), | |
| filename="1_Dense/model.safetensors", | |
| **hf_kwargs, | |
| ) | |
| except Exception: | |
| dense_path = None | |
| if heads_path is not None: | |
| ckpt = torch.load(heads_path, map_location="cpu", weights_only=False) | |
| for k, v in ckpt["projection_heads"].items(): | |
| dim = int(str(k).split(".")[0]) | |
| self._matryoshka_heads[dim] = v.to(self.linear.weight.dtype) | |
| if not self._matryoshka_heads and dense_path is not None: | |
| from safetensors.torch import load_file | |
| sd = load_file(dense_path) | |
| self._matryoshka_heads[ | |
| self.config.embedding_dim | |
| ] = sd["linear.weight"].to(self.linear.weight.dtype) | |
| if not self._matryoshka_heads: | |
| raise FileNotFoundError( | |
| f"Could not load projection weights from {path!r}. " | |
| f"Expected matryoshka_heads.pt or 1_Dense/model.safetensors." | |
| ) | |
| def set_active_dim(self, dim: int) -> None: | |
| """Switch the active Matryoshka embedding dimension.""" | |
| if dim not in self._matryoshka_heads: | |
| available = sorted(self._matryoshka_heads.keys()) | |
| raise ValueError( | |
| f"dim={dim} not available. Available dims: {available}" | |
| ) | |
| w = self._matryoshka_heads[dim] | |
| with torch.no_grad(): | |
| new_linear = nn.Linear(w.size(1), w.size(0), bias=False).to( | |
| device=self.linear.weight.device, dtype=self.linear.weight.dtype | |
| ) | |
| new_linear.weight.copy_(w) | |
| self.linear = new_linear | |
| self._active_dim = dim | |
| self.config.embedding_dim = dim | |
| def active_dim(self) -> int: | |
| return self._active_dim | |
| def matryoshka_dims(self) -> list[int]: | |
| return sorted(self._matryoshka_heads.keys()) | |
| # ---- tokenization helpers ---- | |
| def _resolve_special_ids(self, tokenizer: Any) -> None: | |
| if self._query_prefix_id is None: | |
| q = self.config.query_prefix | |
| for cand in (q, q.strip(), q.rstrip(), q.replace(" ", "")): | |
| tid = tokenizer.convert_tokens_to_ids(cand) | |
| if tid is not None and tid != tokenizer.unk_token_id: | |
| self._query_prefix_id = int(tid) | |
| break | |
| if self._document_prefix_id is None: | |
| d = self.config.document_prefix | |
| for cand in (d, d.strip(), d.rstrip(), d.replace(" ", "")): | |
| tid = tokenizer.convert_tokens_to_ids(cand) | |
| if tid is not None and tid != tokenizer.unk_token_id: | |
| self._document_prefix_id = int(tid) | |
| break | |
| if not self._skiplist_ids: | |
| self._skiplist_ids = [ | |
| tokenizer.convert_tokens_to_ids(w) for w in self.config.skiplist_words | |
| ] | |
| def _tokenize( | |
| self, tokenizer: Any, texts: list[str], is_query: bool | |
| ) -> dict[str, torch.Tensor]: | |
| self._resolve_special_ids(tokenizer) | |
| max_length = ( | |
| self.config.query_length if is_query else self.config.document_length | |
| ) | |
| prefix_id = self._query_prefix_id if is_query else self._document_prefix_id | |
| assert prefix_id is not None, "Failed to resolve ColBERT prefix token id" | |
| do_expand = is_query and self.config.do_query_expansion | |
| # Queries: pad to max_length for ColBERT mask-token expansion. | |
| # Documents: pad to the longest item in the batch so the batch tensor is rectangular. | |
| padding_strategy = "max_length" if do_expand else "longest" | |
| enc = tokenizer( | |
| texts, | |
| padding=padding_strategy, | |
| truncation=True, | |
| max_length=max_length - 1, # reserve 1 slot for prefix token | |
| return_tensors="pt", | |
| ) | |
| input_ids = enc["input_ids"] | |
| attn = enc["attention_mask"] | |
| # Insert prefix at position 1 (right after BOS/CLS). | |
| b = input_ids.size(0) | |
| prefix_col = torch.full( | |
| (b, 1), prefix_id, dtype=input_ids.dtype | |
| ) | |
| attn_col = torch.ones((b, 1), dtype=attn.dtype) | |
| input_ids = torch.cat([input_ids[:, :1], prefix_col, input_ids[:, 1:]], dim=1) | |
| attn = torch.cat([attn[:, :1], attn_col, attn[:, 1:]], dim=1) | |
| if is_query and do_expand and self.config.attend_to_expansion_tokens: | |
| attn = torch.ones_like(attn) | |
| return {"input_ids": input_ids, "attention_mask": attn} | |
| # ---- forward ---- | |
| def forward( | |
| self, | |
| input_ids: torch.Tensor, | |
| attention_mask: torch.Tensor | None = None, | |
| **_: Any, | |
| ) -> BaseModelOutput: | |
| out = self.encoder(input_ids=input_ids, attention_mask=attention_mask) | |
| hidden = out.last_hidden_state if hasattr(out, "last_hidden_state") else out[0] | |
| projected = self.linear(hidden) | |
| normalized = F.normalize(projected, p=2, dim=-1) | |
| return BaseModelOutput(last_hidden_state=normalized) | |
| # ---- public API ---- | |
| def encode( | |
| self, | |
| tokenizer: Any, | |
| texts: list[str], | |
| is_query: bool, | |
| batch_size: int = 16, | |
| ) -> list[torch.Tensor]: | |
| """Return per-text token embeddings (L2-normalized, with skiplist applied).""" | |
| self.eval() | |
| device = next(self.parameters()).device | |
| outputs: list[torch.Tensor] = [] | |
| for start in range(0, len(texts), batch_size): | |
| batch = texts[start : start + batch_size] | |
| feats = self._tokenize(tokenizer, batch, is_query) | |
| feats = {k: v.to(device) for k, v in feats.items()} | |
| out = self( | |
| input_ids=feats["input_ids"], attention_mask=feats["attention_mask"] | |
| ) | |
| token_embs = out.last_hidden_state # (B, L, d) | |
| for i in range(token_embs.size(0)): | |
| ids = feats["input_ids"][i] | |
| mask = feats["attention_mask"][i].bool() | |
| if is_query: | |
| keep = ( | |
| torch.ones_like(ids, dtype=torch.bool) | |
| if self.config.do_query_expansion | |
| else mask | |
| ) | |
| else: | |
| keep = mask.clone() | |
| skiplist = torch.tensor( | |
| self._skiplist_ids, | |
| device=ids.device, | |
| dtype=ids.dtype, | |
| ) | |
| keep &= ~torch.isin(ids, skiplist) | |
| outputs.append(token_embs[i][keep].cpu()) | |
| return outputs | |
| def maxsim(q_emb: torch.Tensor, d_emb: torch.Tensor) -> float: | |
| """ColBERT MaxSim: Σ_q max_d (q · d). Inputs assumed L2-normalized.""" | |
| q = q_emb.to(torch.float32) | |
| d = d_emb.to(torch.float32) | |
| return float((q @ d.T).max(dim=1).values.sum().item()) | |