import spaces # MUST be first — before any torch / CUDA-touching import import os import sys # Make code_modification/ importable so the custom Lfm2BiModel can be loaded sys.path.insert(0, os.path.join(os.path.dirname(__file__), "code_modification")) import torch import gradio as gr # Patch gliner to support Lfm2BiModel before loading the model. # The installed gliner package (v0.2.27) doesn't include LFM2 support. # We register Lfm2BiModel in DECODER_MODEL_MAPPING and patch the # Transformer.__init__ so it doesn't require llm2vec for non-llm2vec configs. import gliner.modeling.encoder as gliner_encoder from lfm2_bi import Lfm2BiModel gliner_encoder.DECODER_MODEL_MAPPING["Lfm2Config"] = Lfm2BiModel # Patch the Transformer.__init__ to use the model repo's logic: # only require IS_LLM2VEC for _LLM2VEC_CONFIGS, not for Lfm2Config _orig_init = gliner_encoder.Transformer.__init__ def _patched_init(self, model_name, config, from_pretrained=False, labels_encoder=False, cache_dir=None): from pathlib import Path from transformers import AutoModel, AutoConfig, DebertaV2Model, T5EncoderModel from gliner.utils import MissedPackageException, is_module_available from gliner.modeling.layers import LayersFuser IS_LLM2VEC = is_module_available("llm2vec") IS_PEFT = is_module_available("peft") if IS_LLM2VEC: from llm2vec.models import GemmaBiModel, LlamaBiModel, Qwen2BiModel, MistralBiModel _EXTRA_MAPPING = { "MistralConfig": MistralBiModel, "LlamaConfig": LlamaBiModel, "GemmaConfig": GemmaBiModel, "Qwen2Config": Qwen2BiModel, } else: _EXTRA_MAPPING = {} import torch.nn as nn nn.Module.__init__(self) if labels_encoder: encoder_config = config.labels_encoder_config else: encoder_config = config.encoder_config if encoder_config is None: encoder_config = AutoConfig.from_pretrained(model_name, cache_dir=cache_dir) if config.vocab_size != -1: encoder_config.vocab_size = config.vocab_size if config._attn_implementation is not None and not labels_encoder: encoder_config._attn_implementation = config._attn_implementation config_name = encoder_config.__class__.__name__ kwargs = {} _LLM2VEC_CONFIGS = {"MistralConfig", "LlamaConfig", "GemmaConfig", "Qwen2Config"} # Build the full mapping: base package's + our LFM2 addition + llm2vec if available FULL_MAPPING = dict(gliner_encoder.DECODER_MODEL_MAPPING) FULL_MAPPING.update(_EXTRA_MAPPING) if config_name in FULL_MAPPING: if config_name in _LLM2VEC_CONFIGS and not IS_LLM2VEC: raise MissedPackageException( f"The llm2vec package must be installed to use this decoder model: {config_name}" ) ModelClass = FULL_MAPPING[config_name] custom = True elif config_name in {"T5Config", "MT5Config"}: custom = True ModelClass = T5EncoderModel elif config_name in {"DebertaV2Config"}: custom = True ModelClass = DebertaV2Model else: custom = False ModelClass = AutoModel if from_pretrained: self.model = ModelClass.from_pretrained(model_name, **kwargs, trust_remote_code=True) elif not custom: self.model = ModelClass.from_config(encoder_config, trust_remote_code=True) else: self.model = ModelClass(encoder_config, **kwargs) adapter_config_file = Path(model_name) / "adapter_config.json" if adapter_config_file.exists(): if IS_PEFT: from peft import LoraConfig, get_peft_model adapter_config = LoraConfig.from_pretrained(model_name) self.model = get_peft_model(self.model, adapter_config) else: import warnings warnings.warn( "Adapter configs were detected, if you want to apply them you need to install peft package.", stacklevel=2, ) if config.fuse_layers: self.layers_fuser = LayersFuser(encoder_config.num_hidden_layers, encoder_config.hidden_size) if labels_encoder: config.labels_encoder_config = encoder_config else: config.encoder_config = encoder_config self.config = config gliner_encoder.Transformer.__init__ = _patched_init from gliner import GLiNER MODEL_ID = "VAGOsolutions/SauerkrautLM-LFM2.5-GLiNER" # Load at module scope, .to("cuda") eagerly — ZeroGPU intercepts and packs model = GLiNER.from_pretrained(MODEL_ID).to("cuda") EXAMPLES = [ [ "Maria Schmidt arbeitet bei Siemens in München, E-Mail: maria.schmidt@siemens.com", "person, organization, location, email", 0.5, False, ], [ "John Doe called from +1-202-555-0173 to report that his credit card " "4111-1111-1111-1111 expires on 03/2027.", "person, phone number, credit card number, date", 0.5, False, ], [ "Le 15 mars 2025, la société Renault a annoncé un partenariat avec " "l'Université Paris-Sorbonne pour développer des véhicules autonomes.", "personne, organisation, lieu, date", 0.5, False, ], [ "Der Patient Hans Müller wurde am 14.02.2024 im Universitätsklinikum " "Heidelberg mit der Diagnose Diabetes mellitus Typ 2 aufgenommen. " "Seine Krankenversicherungsnummer lautet A123456789B.", "person, krankenhaus, datum, diagnose, versicherungsnummer", 0.5, False, ], [ "On July 20, 1969, Neil Armstrong and Buzz Aldrin landed the Apollo 11 " "lunar module Eagle at the Sea of Tranquility. The mission was operated " "by NASA from the Manned Spacecraft Center in Houston, Texas.", "person, organization, date, location, vehicle", 0.5, True, ], [ "La Sra. García trabaja en Telefónica desde 2018 y vive en el barrio " "de Salamanca, Madrid. Su correo es m.garcia@telefonica.com.", "persona, organización, ubicación, fecha, correo electrónico", 0.5, False, ], ] @spaces.GPU(duration=30) def ner( text: str, labels: str, threshold: float, nested_ner: bool, ): """Extract named entities from text using zero-shot GLiNER. Provide any entity types as comma-separated labels — the model extracts matching spans without retraining. Args: text: The input text to analyze. labels: Comma-separated entity types (e.g. "person, location, date"). threshold: Confidence threshold (lower = more entities, higher = fewer). nested_ner: If True, allow nested/overlapping entity spans. """ label_list = [l.strip() for l in labels.split(",") if l.strip()] entities = model.predict_entities( text, label_list, flat_ner=not nested_ner, threshold=threshold ) return { "text": text, "entities": [ { "entity": ent["label"], "word": ent["text"], "start": ent["start"], "end": ent["end"], "score": ent.get("score", 0.0), } for ent in entities ], } CSS = """ #col-container { max-width: 1100px; margin: 0 auto; } .dark .gradio-container { color: var(--body-text-color); } """ with gr.Blocks(title="SauerkrautLM-LFM2.5-GLiNER") as demo: gr.Markdown( """ # SauerkrautLM-LFM2.5-GLiNER — Zero-Shot NER A compact **350M** bidirectional LFM2.5 model for **zero-shot Named Entity Recognition**. Enter any text and a comma-separated list of entity types — the model extracts matching spans without retraining. Supports **English, French, German, Italian, and Spanish**. Strong on general NER, PII/privacy, and biomedical entities. ## Links * Model: [VAGOsolutions/SauerkrautLM-LFM2.5-GLiNER](https://huggingface.co/VAGOsolutions/SauerkrautLM-LFM2.5-GLiNER) * GLiNER library: [github.com/urchade/GLiNER](https://github.com/urchade/GLiNER) * Paper: [arxiv.org/abs/2311.08526](https://arxiv.org/abs/2311.08526) """ ) with gr.Column(elem_id="col-container"): input_text = gr.Textbox( value=EXAMPLES[0][0], label="Text input", placeholder="Enter your text here", lines=6, ) with gr.Row(): labels = gr.Textbox( value=EXAMPLES[0][1], label="Labels", placeholder="person, location, date, ...", scale=2, ) threshold = gr.Slider( 0, 1, value=0.5, step=0.01, label="Threshold", info="Lower = more entities; higher = more precise.", scale=1, ) nested_ner = gr.Checkbox( value=False, label="Nested NER", info="Allow overlapping entity spans.", scale=0, ) output = gr.HighlightedText(label="Predicted Entities") submit_btn = gr.Button("Extract entities", variant="primary") gr.Examples( examples=EXAMPLES, fn=ner, inputs=[input_text, labels, threshold, nested_ner], outputs=output, cache_examples=True, cache_mode="lazy", ) # Wire up events input_text.submit(fn=ner, inputs=[input_text, labels, threshold, nested_ner], outputs=output) labels.submit(fn=ner, inputs=[input_text, labels, threshold, nested_ner], outputs=output) threshold.release(fn=ner, inputs=[input_text, labels, threshold, nested_ner], outputs=output) nested_ner.change(fn=ner, inputs=[input_text, labels, threshold, nested_ner], outputs=output) submit_btn.click(fn=ner, inputs=[input_text, labels, threshold, nested_ner], outputs=output, api_name="ner") demo.queue() demo.launch(mcp_server=True, theme=gr.themes.Citrus(), css=CSS)