Views
No views yet
CC(=O)O). The [SEP] token (id=13) is used as the end-of-sequence marker.1from transformers import AutoModelForMaskedLM, AutoTokenizer
2import torch
3
4model = AutoModelForMaskedLM.from_pretrained("kojima-lab/molcrawl-compounds-chemberta2-medium")
5tokenizer = AutoTokenizer.from_pretrained("kojima-lab/molcrawl-compounds-chemberta2-medium")
6
7# Predict masked SMILES 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 = "CC(=O){MASK}".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_compounds_chemberta2_medium,
2 title={molcrawl-compounds-chemberta2-medium},
3 author={{RIKEN}},
4 year={2026},
5 publisher={{Hugging Face}},
6 url={{https://huggingface.co/kojima-lab/molcrawl-compounds-chemberta2-medium}}
7}1import torch
2from transformers import AutoTokenizer, AutoModelForMaskedLM
3
4REPO_ID = "kojima-lab/molcrawl-compounds-chemberta2-medium"
5tokenizer = AutoTokenizer.from_pretrained(REPO_ID)
6model = AutoModelForMaskedLM.from_pretrained(REPO_ID)
7model.eval()
8
9# SMILES with one masked position
10prompt = "CC(=O)Oc1ccccc1[MASK](=O)O"
11inputs = tokenizer(prompt, return_tensors="pt")
12mask_index = (inputs["input_ids"][0] == tokenizer.mask_token_id).nonzero(as_tuple=True)[0]
13
14with torch.no_grad():
15 outputs = model(**inputs)
16
17predicted_id = outputs.logits[0, mask_index].argmax(dim=-1)
18predicted_token = tokenizer.convert_ids_to_tokens(predicted_id.tolist())[0]
19print(f"Predicted token at mask: {predicted_token}")
20# => Predicted token at mask: C