Views
No views yet
1import torch
2import numpy as np
3from transformers import AutoModel, AutoTokenizer
4import torch.nn.functional as F
5
6MODEL_ID = "THUMedInfo/DR.EHR-small"
7device = "cuda" if torch.cuda.is_available() else "cpu"
8max_length = 512 # note chunks
9batch_size = 32
10
11tokenizer = AutoTokenizer.from_pretrained(MODEL_ID, trust_remote_code=True)
12model = AutoModel.from_pretrained(MODEL_ID, trust_remote_code=True).to(device)
13model.eval()
14
15@torch.no_grad()
16def embed_texts(texts):
17 all_emb = []
18 for i in range(0, len(texts), batch_size):
19 batch = texts[i:i+batch_size]
20 enc = tokenizer(
21 batch,
22 padding=True,
23 truncation=True,
24 max_length=max_length,
25 return_tensors="pt",
26 return_token_type_ids=False,
27 )
28 enc = {k: v.to(device) for k, v in enc.items()}
29 out = model(**enc)
30
31 # CLS pooling (BERT-style)
32 emb = out.last_hidden_state[:, 0, :] # [B, 768]
33 emb = F.normalize(emb, p=2, dim=1)
34 all_emb.append(emb.cpu().numpy())
35 return np.vstack(all_emb)
36
37# Example
38queries = ["hypertension", "metformin"]
39q_emb = embed_texts(queries)
40print(q_emb.shape)@article{zhao2025dr,
title={DR. EHR: Dense Retrieval for Electronic Health Record with Knowledge Injection and Synthetic Data},
author={Zhao, Zhengyun and Ying, Huaiyuan and Zhong, Yue and Yu, Sheng},
journal={arXiv preprint arXiv:2507.18583},
year={2025}
}