This model is trained through the approach described in DMRetriever: A Family of Models for Improved Text Retrieval in Disaster Management.
The associated GitHub repository is available here.
This model has 596M parameters and it is the pre-trained version (trained using only unlabeled dataset containing in-batch negative).
🧠 Model Overview
DMRetriever-596M-PT has the following features:
- Model Type: Text Embedding
- Supported Languages: English
- Number of Paramaters: 0.6B
- Embedding Dimension: 1024
For more details, including model training, benchmark evaluation, and inference performance, please refer to our paper, GitHub.
📦 DMRetriever Series Model List
🚀 Usage
Using HuggingFace Transformers:
import torch
import torch.nn.functional as F
from transformers import AutoTokenizer
from bidirectional_qwen3 import Qwen3BiModel
MODEL_ID = "DMIR01/DMRetriever-596M-PT"
device = "cuda" if torch.cuda.is_available() else "cpu"
dtype = torch.float16 if device == "cuda" else torch.float32
tokenizer = AutoTokenizer.from_pretrained(
MODEL_ID,
trust_remote_code=True,
use_fast=False,
)
if getattr(tokenizer, "pad_token_id", None) is None and getattr(tokenizer, "eos_token", None) is not None:
tokenizer.pad_token = tokenizer.eos_token
tokenizer.padding_side = "right"
model = Qwen3BiModel.from_pretrained(
MODEL_ID,
torch_dtype=dtype,
trust_remote_code=True,
).to(device).eval()
def mean_pool(last_hidden_state: torch.Tensor, attention_mask: torch.Tensor) -> torch.Tensor:
mask = attention_mask.unsqueeze(-1).to(last_hidden_state.dtype)
summed = (last_hidden_state * mask).sum(dim=1)
counts = mask.sum(dim=1).clamp(min=1e-9)
return summed / counts
def encode_texts(texts, batch_size=32, max_length=512):
vecs = []
for i in range(0, len(texts), batch_size):
batch = texts[i:i+batch_size]
with torch.no_grad():
inputs = tokenizer(
batch,
max_length=max_length,
truncation=True,
padding=True,
return_tensors="pt",
).to(device)
hidden = model(**inputs).last_hidden_state
emb = mean_pool(hidden, inputs["attention_mask"])
emb = F.normalize(emb, p=2, dim=1)
vecs.append(emb.cpu())
return torch.cat(vecs, dim=0) if vecs else torch.empty(0, model.config.hidden_size)
TASK2PREFIX = {
"FactCheck": "Given the claim, retrieve most relevant document that supports or refutes the claim",
"NLI": "Given the premise, retrieve most relevant hypothesis that is entailed by the premise",
"QA": "Given the question, retrieve most relevant passage that best answers the question",
"QAdoc": "Given the question, retrieve the most relevant document that answers the question",
"STS": "Given the sentence, retrieve the sentence with the same meaning",
"Twitter": "Given the user query, retrieve the most relevant Twitter text that meets the request",
}
def apply_task_prefix(queries, task: str):
"""Add instruction to queries; corpus texts remain unchanged."""
prefix = TASK2PREFIX.get(task, "")
if prefix:
return [f"{prefix}: {q.strip()}" for q in queries]
return [q.strip() for q in queries]
task = "QA"
queries_raw = [
"Who wrote The Little Prince?",
"What is the capital of France?",
]
queries = apply_task_prefix(queries_raw, task)
corpus_passages = [
"The Little Prince is a novella by Antoine de Saint-Exupéry, first published in 1943.",
"Paris is the capital and most populous city of France.",
"Transformers are neural architectures that rely on attention mechanisms.",
]
query_emb = encode_texts(queries, batch_size=32, max_length=512)
corpus_emb = encode_texts(corpus_passages, batch_size=32, max_length=512)
print("Query embeddings:", tuple(query_emb.shape))
print("Corpus embeddings:", tuple(corpus_emb.shape))
scores = query_emb @ corpus_emb.T
topk = scores.topk(k=min(3, corpus_emb.size(0)), dim=1)
for i, q in enumerate(queries_raw):
print(f"\nQuery[{i}] {q}")
for rank, (score, idx) in enumerate(zip(topk.values[i].tolist(), topk.indices[i].tolist()), start=1):
print(f" Top{rank}: doc#{idx} | score={score:.4f} | text={corpus_passages[idx]}")
⚠️ Notice
The backbone used in DMRetriever is Bidirectional Qwen3, not the standard Qwen3.
Please ensure that the bidirectional_qwen3 module (included in the released model checkpoint folder) is correctly placed inside your model directory.
Make sure that your transformers library version is > 4.51.0 to avoid the error:
KeyError: 'qwen3'.
🧾 Citation
If you find this repository helpful, please kindly consider citing the corresponding paper. Thanks!
@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}
}