Views
No views yet
molcrawl-molecule-nat-lang-gpt2-medium pre-trained model.1from transformers import AutoModelForMaskedLM, AutoTokenizer
2import torch
3
4model = AutoModelForMaskedLM.from_pretrained("kojima-lab/molcrawl-molecule-nat-lang-mol-instructions-bert-medium")
5tokenizer = AutoTokenizer.from_pretrained("kojima-lab/molcrawl-molecule-nat-lang-mol-instructions-bert-medium")
6
7# Predict masked token
8# Use tokenizer.mask_token instead of hardcoded "[MASK]":
9# BERT-style tokenizers vary ("[MASK]", "<mask>", etc.)
10if tokenizer.mask_token is None:
11 raise ValueError("This tokenizer has no mask_token; masked LM inference is not supported.")
12prompt = "your input {MASK} sequence".replace("{MASK}", tokenizer.mask_token)
13inputs = tokenizer(prompt, return_tensors="pt")
14mask_index = (inputs["input_ids"] == tokenizer.mask_token_id).nonzero(as_tuple=True)[1]
15
16with torch.no_grad():
17 outputs = model(**inputs)
18logits = outputs.logits
19
20predicted_token_id = logits[0, mask_index].argmax(dim=-1)
21predicted_token = tokenizer.decode(predicted_token_id)
22result = prompt.replace(tokenizer.mask_token, predicted_token)
23print(f"Predicted: {result}")
241@misc{molcrawl_molecule_nat_lang_mol_instructions_bert_medium,
2 title={molcrawl-molecule-nat-lang-mol-instructions-bert-medium},
3 author={{RIKEN}},
4 year={2026},
5 publisher={{Hugging Face}},
6 url={{https://huggingface.co/kojima-lab/molcrawl-molecule-nat-lang-mol-instructions-bert-medium}}
7}1import torch
2from transformers import AutoTokenizer, AutoModelForMaskedLM
3
4REPO_ID = "kojima-lab/molcrawl-molecule-nat-lang-mol-instructions-bert-medium"
5tokenizer = AutoTokenizer.from_pretrained(REPO_ID)
6model = AutoModelForMaskedLM.from_pretrained(REPO_ID)
7model.eval()
8
9def embed(text):
10 inputs = tokenizer(text, return_tensors="pt", truncation=True, max_length=128)
11 with torch.no_grad():
12 out = model.bert(**inputs) # encoder only; pooler is unused / random-init
13 mask = inputs["attention_mask"][0].unsqueeze(-1).float()
14 return ((out.last_hidden_state[0] * mask).sum(0) / mask.sum())
15
16texts = [
17 "Aspirin is an anti-inflammatory drug.",
18 "Ibuprofen is an anti-inflammatory drug.",
19 "Glucose is a simple sugar.",
20 "DNA stores genetic information.",
21]
22embs = torch.stack([embed(t) for t in texts])
23embs = embs / embs.norm(dim=-1, keepdim=True)
24
25for i in range(len(texts)):
26 for j in range(i + 1, len(texts)):
27 print(f"sim('{texts[i][:24]}...', '{texts[j][:24]}...') = {(embs[i] @ embs[j]).item():.3f}")
28# Expected (approximately):
29# sim('Aspirin is an anti-infl...', 'Ibuprofen is an anti-inf...') = 0.980
30# sim('Aspirin is an anti-infl...', 'Glucose is a simple suga...') = 0.964
31# sim('Aspirin is an anti-infl...', 'DNA stores genetic infor...') = 0.921
32# sim('Ibuprofen is an anti-in...', 'Glucose is a simple suga...') = 0.947
33# sim('Ibuprofen is an anti-in...', 'DNA stores genetic infor...') = 0.885
34# sim('Glucose is a simple sug...', 'DNA stores genetic infor...') = 0.957