Views
No views yet
| Model | Description | Backbone | Backbone Type | Hidden Size | #Layers |
|---|---|---|---|---|---|
| DMRetriever-33M | Base 33M variant | MiniLM | Encoder-only | 384 | 12 |
| DMRetriever-33M-PT | Pre-trained version of 33M | MiniLM | Encoder-only | 384 | 12 |
| DMRetriever-109M | Base 109M variant | BERT-base-uncased | Encoder-only | 768 | 12 |
| DMRetriever-109M-PT | Pre-trained version of 109M | BERT-base-uncased | Encoder-only | 768 | 12 |
| DMRetriever-335M | Base 335M variant | BERT-large-uncased-WWM | Encoder-only | 1024 | 24 |
| DMRetriever-335M-PT | Pre-trained version of 335M | BERT-large-uncased-WWM | Encoder-only | 1024 | 24 |
| DMRetriever-596M | Base 596M variant | Qwen3-0.6B | Decoder-only | 1024 | 28 |
| DMRetriever-596M-PT | Pre-trained version of 596M | Qwen3-0.6B | Decoder-only | 1024 | 28 |
| DMRetriever-4B | Base 4B variant | Qwen3-4B | Decoder-only | 2560 | 36 |
| DMRetriever-4B-PT | Pre-trained version of 4B | Qwen3-4B | Decoder-only | 2560 | 36 |
| DMRetriever-7.6B | Base 7.6B variant | Qwen3-8B | Decoder-only | 4096 | 36 |
| DMRetriever-7.6B-PT | Pre-trained version of 7.6B | Qwen3-8B | Decoder-only | 4096 | 36 |
1# pip install torch transformers
2import torch
3import torch.nn.functional as F
4from transformers import AutoTokenizer
5from bidirectional_qwen3 import Qwen3BiModel # custom bidirectional backbone
6
7MODEL_ID = "DMIR01/DMRetriever-4B"
8
9# Device & dtype
10device = "cuda" if torch.cuda.is_available() else "cpu"
11dtype = torch.float16 if device == "cuda" else torch.float32
12
13# --- Tokenizer (needs remote code for custom modules) ---
14tokenizer = AutoTokenizer.from_pretrained(
15 MODEL_ID,
16 trust_remote_code=True,
17 use_fast=False,
18)
19# Ensure pad token and right padding (matches training)
20if getattr(tokenizer, "pad_token_id", None) is None and getattr(tokenizer, "eos_token", None) is not None:
21 tokenizer.pad_token = tokenizer.eos_token
22tokenizer.padding_side = "right"
23
24# --- Bidirectional encoder (non-autoregressive; for retrieval/embedding) ---
25model = Qwen3BiModel.from_pretrained(
26 MODEL_ID,
27 torch_dtype=dtype,
28 trust_remote_code=True,
29).to(device).eval()
30
31# --- Mean pooling over valid tokens ---
32def mean_pool(last_hidden_state: torch.Tensor, attention_mask: torch.Tensor) -> torch.Tensor:
33 mask = attention_mask.unsqueeze(-1).to(last_hidden_state.dtype) # [B, L, 1]
34 summed = (last_hidden_state * mask).sum(dim=1) # [B, H]
35 counts = mask.sum(dim=1).clamp(min=1e-9) # [B, 1]
36 return summed / counts
37
38# --- Batch encoder: returns L2-normalized embeddings ---
39def encode_texts(texts, batch_size=32, max_length=512):
40 vecs = []
41 for i in range(0, len(texts), batch_size):
42 batch = texts[i:i+batch_size]
43 with torch.no_grad():
44 inputs = tokenizer(
45 batch,
46 max_length=max_length,
47 truncation=True,
48 padding=True,
49 return_tensors="pt",
50 ).to(device)
51 hidden = model(**inputs).last_hidden_state
52 emb = mean_pool(hidden, inputs["attention_mask"])
53 emb = F.normalize(emb, p=2, dim=1) # cosine-ready
54 vecs.append(emb.cpu())
55 return torch.cat(vecs, dim=0) if vecs else torch.empty(0, model.config.hidden_size)
56
57# --- Task instructions (apply to queries only) ---
58TASK2PREFIX = {
59 "FactCheck": "Given the claim, retrieve most relevant document that supports or refutes the claim",
60 "NLI": "Given the premise, retrieve most relevant hypothesis that is entailed by the premise",
61 "QA": "Given the question, retrieve most relevant passage that best answers the question",
62 "QAdoc": "Given the question, retrieve the most relevant document that answers the question",
63 "STS": "Given the sentence, retrieve the sentence with the same meaning",
64 "Twitter": "Given the user query, retrieve the most relevant Twitter text that meets the request",
65}
66
67def apply_task_prefix(queries, task: str):
68 """Add instruction to queries; corpus texts remain unchanged."""
69 prefix = TASK2PREFIX.get(task, "")
70 if prefix:
71 return [f"{prefix}: {q.strip()}" for q in queries]
72 return [q.strip() for q in queries]
73
74# ========================= Usage =========================
75# Queries need task instruction
76task = "QA"
77queries_raw = [
78 "Who wrote The Little Prince?",
79 "What is the capital of France?",
80]
81queries = apply_task_prefix(queries_raw, task)
82
83# Corpus: no instruction
84corpus_passages = [
85 "The Little Prince is a novella by Antoine de Saint-Exupéry, first published in 1943.",
86 "Paris is the capital and most populous city of France.",
87 "Transformers are neural architectures that rely on attention mechanisms.",
88]
89
90# Encode
91query_emb = encode_texts(queries, batch_size=32, max_length=512) # [Q, H]
92corpus_emb = encode_texts(corpus_passages, batch_size=32, max_length=512) # [D, H]
93print("Query embeddings:", tuple(query_emb.shape))
94print("Corpus embeddings:", tuple(corpus_emb.shape))
95
96# Retrieval demo: cosine similarity via dot product (embeddings are normalized)
97scores = query_emb @ corpus_emb.T # [Q, D]
98topk = scores.topk(k=min(3, corpus_emb.size(0)), dim=1)
99
100for i, q in enumerate(queries_raw):
101 print(f"\nQuery[{i}] {q}")
102 for rank, (score, idx) in enumerate(zip(topk.values[i].tolist(), topk.indices[i].tolist()), start=1):
103 print(f" Top{rank}: doc#{idx} | score={score:.4f} | text={corpus_passages[idx]}")
104
105bidirectional_qwen3 module (included in the released model checkpoint folder) is correctly placed inside your model directory.KeyError: 'qwen3'.@article{yin2025dmretriever,
title={DMRetriever: A Family of Models for Improved Text Retrieval in Disaster Management},
author={Yin, Kai and Dong, Xiangjue and Liu, Chengkai and Lin, Allen and Shi, Lingfeng and Mostafavi, Ali and Caverlee, James},
journal={arXiv preprint arXiv:2510.15087},
year={2025}
}