Fill-Mask
Transformers
English
masked-language-modeling
summarization-evaluation
entity-infilling
mars
modernbert
Instructions to use Glazkov/mars-shared-cross-attention-modernbert with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use Glazkov/mars-shared-cross-attention-modernbert with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("fill-mask", model="Glazkov/mars-shared-cross-attention-modernbert")# Load model directly from transformers import AutoModel model = AutoModel.from_pretrained("Glazkov/mars-shared-cross-attention-modernbert", device_map="auto") - Notebooks
- Google Colab
- Kaggle
| """High-level inference helper for the MARS shared-cross-attention checkpoint. | |
| Uses parallel multi-mask decoding: one forward pass per token-in-span across | |
| all masks simultaneously. Matches the recipe used to score MARSv2 = 57.02 | |
| on the CNN/DailyMail benchmark (best single-model result in the project). | |
| """ | |
| from typing import Optional | |
| import torch | |
| import torch.nn.functional as F | |
| from modeling_mars import ( | |
| ENTEND, | |
| ENTMASK, | |
| ENTSTART, | |
| TYPED_ENTMASK_TOKENS, | |
| load_model_from_checkpoint, | |
| ) | |
| class MarsInference: | |
| def __init__( | |
| self, | |
| repo_or_path: str, | |
| base_model: str = "answerdotai/ModernBERT-base", | |
| device: Optional[str] = None, | |
| ): | |
| self.model, self.tokenizer, self.device = load_model_from_checkpoint( | |
| repo_or_path, base_model=base_model, device=device | |
| ) | |
| self.entmask_id = self.tokenizer.convert_tokens_to_ids(ENTMASK) | |
| self.entstart_id = self.tokenizer.convert_tokens_to_ids(ENTSTART) | |
| self.entend_id = self.tokenizer.convert_tokens_to_ids(ENTEND) | |
| self.typed_entmask_ids = { | |
| t: self.tokenizer.convert_tokens_to_ids(tok) | |
| for t, tok in TYPED_ENTMASK_TOKENS.items() | |
| } | |
| def _replace_masks(self, masked_text: str, entity_types: Optional[list[str]]) -> str: | |
| if not entity_types: | |
| return masked_text.replace("<mask>", f"{ENTSTART} {ENTMASK}") | |
| parts = masked_text.split("<mask>") | |
| out = parts[0] | |
| for j in range(len(parts) - 1): | |
| tok = TYPED_ENTMASK_TOKENS.get(entity_types[j], ENTMASK) if j < len(entity_types) else ENTMASK | |
| out += f"{ENTSTART} {tok}" + parts[j + 1] | |
| return out | |
| def _is_mask(self, ids: torch.Tensor) -> torch.Tensor: | |
| m = ids == self.entmask_id | |
| for tid in self.typed_entmask_ids.values(): | |
| m = m | (ids == tid) | |
| return m | |
| def predict( | |
| self, | |
| summary: str, | |
| masked_text: str, | |
| max_length: int = 512, | |
| max_span_len: int = 8, | |
| entity_types: Optional[list[str]] = None, | |
| return_confidence: bool = False, | |
| ): | |
| """Predict the entities that fill each `<mask>` in `masked_text`. | |
| Args: | |
| summary: short summary text providing the grounding context. | |
| masked_text: original text with each entity replaced by `<mask>`. | |
| entity_types: optional spaCy NER types per mask (e.g. "PERSON", | |
| "ORG", ...). When given, the model uses the matching typed | |
| mask token which improves accuracy. | |
| Returns: list of predicted entity strings, one per `<mask>` (and | |
| optionally a parallel list of mean-token confidences). | |
| """ | |
| processed = self._replace_masks(masked_text, entity_types) | |
| summary_enc = self.tokenizer( | |
| summary, padding=True, truncation=True, | |
| max_length=max_length // 2, return_tensors="pt", | |
| ).to(self.device) | |
| m_tok = self.tokenizer(processed, add_special_tokens=False, return_tensors="pt") | |
| masked_ids = m_tok["input_ids"].squeeze(0) | |
| end_ids = {self.entend_id, self.tokenizer.sep_token_id} | |
| if self.tokenizer.eos_token_id is not None: | |
| end_ids.add(self.tokenizer.eos_token_id) | |
| positions = self._is_mask(masked_ids).nonzero(as_tuple=False).squeeze(-1).tolist() | |
| n = len(positions) | |
| if n == 0: | |
| return ([], []) if return_confidence else [] | |
| spans: list[list[int]] = [[] for _ in range(n)] | |
| confs: list[list[float]] = [[] for _ in range(n)] | |
| done: list[bool] = [False] * n | |
| for _ in range(max_span_len): | |
| active = [(i, positions[i]) for i in range(n) if not done[i]] | |
| if not active: | |
| break | |
| inp = masked_ids.unsqueeze(0).to(self.device) | |
| attn = torch.ones_like(inp) | |
| out = self.model( | |
| summary_input_ids=summary_enc["input_ids"], | |
| summary_attention_mask=summary_enc["attention_mask"], | |
| masked_input_ids=inp, | |
| masked_attention_mask=attn, | |
| ) | |
| logits = out.logits | |
| insertions = [] | |
| for ent_idx, pos in active: | |
| tl = logits[0, pos, :].clone() | |
| tl[self.entmask_id] = float("-inf") | |
| tl[self.entstart_id] = float("-inf") | |
| for tid in self.typed_entmask_ids.values(): | |
| tl[tid] = float("-inf") | |
| p = F.softmax(tl, dim=-1) | |
| nid = int(torch.argmax(p).item()) | |
| c = float(p[nid].item()) | |
| if nid in end_ids: | |
| done[ent_idx] = True | |
| else: | |
| insertions.append((ent_idx, pos, nid, c)) | |
| if not insertions: | |
| break | |
| for ent_idx, pos, tok_id, c in sorted(insertions, key=lambda x: x[1], reverse=True): | |
| new_tok = torch.tensor([tok_id], dtype=masked_ids.dtype) | |
| masked_ids = torch.cat([masked_ids[:pos], new_tok, masked_ids[pos:]]) | |
| spans[ent_idx].append(tok_id) | |
| confs[ent_idx].append(c) | |
| positions[ent_idx] = pos + 1 | |
| for j in range(n): | |
| if j != ent_idx and positions[j] >= pos: | |
| positions[j] += 1 | |
| texts = [ | |
| self.tokenizer.decode(s, skip_special_tokens=True, clean_up_tokenization_spaces=True).strip() | |
| for s in spans | |
| ] | |
| if return_confidence: | |
| avg = [sum(c) / max(1, len(c)) for c in confs] | |
| return texts, avg | |
| return texts | |
| if __name__ == "__main__": | |
| import sys | |
| repo = sys.argv[1] if len(sys.argv) > 1 else "Glazkov/mars-shared-cross-attention-modernbert" | |
| summary = ( | |
| "The president announced a new climate policy in Washington on Tuesday, " | |
| "promising to cut emissions by 40% by 2030." | |
| ) | |
| masked = ( | |
| "<mask> announced a new climate policy in <mask> on <mask>, " | |
| "promising to cut emissions by <mask> by <mask>." | |
| ) | |
| types = ["PERSON", "GPE", "DATE", "PERCENT", "DATE"] | |
| inf = MarsInference(repo) | |
| preds, c = inf.predict(summary, masked, entity_types=types, return_confidence=True) | |
| for m, p, x in zip(types, preds, c): | |
| print(f" [{m}] -> {p!r} (conf={x:.2f})") | |