Views
No views yet
answerdotai/ModernBERT-basetorch.bfloat16 for efficient computation on modern hardwaretransformers library. Until the next transformers release, doing so requires installing transformers from main:pip install git+https://github.com/huggingface/transformers.gitfill-mask pipeline or load it via AutoModelForMaskedLM. To use ModernBERT for downstream tasks like classification, retrieval, or QA, fine-tune it following standard BERT fine-tuning recipes.pip install flash-attn1import torch
2from transformers import AutoTokenizer, AutoModelForMaskedLM
3
4# Load custom tokenizer and fine-tuned model
5tokenizer = AutoTokenizer.from_pretrained("myrkur/Persian-ModernBert-base")
6model = AutoModelForMaskedLM.from_pretrained("myrkur/Persian-ModernBert-base", attn_implementation="eager", torch_dtype=torch.bfloat16, device_map="cpu")1text = "حال و [MASK] مردم خوب است."
2inputs = tokenizer(text, return_tensors="pt")
3inputs = {k:v.cpu() for k, v in inputs.items()}
4token_logits = model(**inputs).logits
5
6# Find the [MASK] token and decode top predictions
7mask_token_index = torch.where(inputs["input_ids"] == tokenizer.mask_token_id)[1]
8mask_token_logits = token_logits[0, mask_token_index, :]
9top_5_tokens = torch.topk(mask_token_logits, 5, dim=1).indices[0].tolist()
10
11for token in top_5_tokens:
12 print(f"Prediction: {text.replace(tokenizer.mask_token, tokenizer.decode([token]))}")1import torch
2from transformers import AutoTokenizer, AutoModelForMaskedLM
3
4# Load custom tokenizer and fine-tuned model
5tokenizer = AutoTokenizer.from_pretrained("myrkur/Persian-ModernBert-base")
6model = AutoModelForMaskedLM.from_pretrained("myrkur/Persian-ModernBert-base", attn_implementation="flash_attention_2", torch_dtype=torch.bfloat16, device_map="cuda")1text = "حال و [MASK] مردم خوب است."
2inputs = tokenizer(text, return_tensors="pt")
3inputs = {k:v.cuda() for k, v in inputs.items()}
4token_logits = model(**inputs).logits
5
6# Find the [MASK] token and decode top predictions
7mask_token_index = torch.where(inputs["input_ids"] == tokenizer.mask_token_id)[1]
8mask_token_logits = token_logits[0, mask_token_index, :]
9top_5_tokens = torch.topk(mask_token_logits, 5, dim=1).indices[0].tolist()
10
11for token in top_5_tokens:
12 print(f"Prediction: {text.replace(tokenizer.mask_token, tokenizer.decode([token]))}")flash_attention_2 implementation, significantly reducing memory overhead while accelerating training on large datasets.