CoolFace
Datasetpublic

baconnier/xp-ft-workspace-fr-016c

ml-models/ — Modeles ML en production Derniere mise a jour : 28 mars 2026 Monte dans le container ml-light-spark sous /models/. Modeles en production Modele Taille Usage Charge par gliner-notarial-v4 1.7G NER (entites nommees) GLINER_MODEL env var gliclass-notarial-v2 1.5G Classification documents (45 classes) GLICLASS_MODEL env var bge-m3-notarial-v3 2.2G Embeddings (dense + sparse + ColBERT) BGE_M3_MODEL env var gliner-relex-notarial-v1 3.6G… See the full description on the dataset page: https://huggingface.co/datasets/baconnier/xp-ft-workspace-fr-016c.

sourceHugging Faceupdated 2mo agoView on Hugging Face
0likes95downloads
gliner_finetune.py546 linesDownload Raw Back to root
1#!/usr/bin/env python32"""GLiNER Fine-Tuning for Notarial Documents (Phase 3).3 4Fine-tunes gliner-multitask-large-v0.5 on annotated notarial entity data.5Training args are tuned for DGX Spark (1 GPU, fp16).6 7Usage:8    python3 scripts/gliner_finetune.py \9        --dataset-dir data/gliner \10        --output-dir models/gliner-notarial-v111 12    # Dry run (load data, print stats, no training):13    python3 scripts/gliner_finetune.py \14        --dataset-dir data/gliner \15        --output-dir models/gliner-notarial-v1 \16        --dry-run17"""18 19from __future__ import annotations20 21import argparse22import json23import logging24import os25import sys26from collections import Counter27from pathlib import Path28from typing import Any, Dict, List, Optional29 30logging.basicConfig(31    level=logging.INFO,32    format="%(asctime)s [%(levelname)s] %(message)s",33    datefmt="%Y-%m-%d %H:%M:%S",34)35logger = logging.getLogger("gliner_finetune")36 37# ---------------------------------------------------------------------------38# Dependency checks39# ---------------------------------------------------------------------------40 41def _check_gliner_available() -> bool:42    try:43        import gliner  # noqa: F40144        return True45    except ImportError:46        return False47 48 49def _check_training_available() -> bool:50    try:51        from gliner.training import Trainer, TrainingArguments  # noqa: F40152        return True53    except (ImportError, AttributeError):54        return False55 56 57# ---------------------------------------------------------------------------58# Dataset loading59# ---------------------------------------------------------------------------60 61def load_jsonl(path: Path) -> List[Dict[str, Any]]:62    """Load a JSONL file, one JSON object per line."""63    samples: List[Dict[str, Any]] = []64    with open(path, "r", encoding="utf-8") as f:65        for i, line in enumerate(f, 1):66            line = line.strip()67            if not line:68                continue69            try:70                samples.append(json.loads(line))71            except json.JSONDecodeError as exc:72                logger.warning("Skipping malformed line %d in %s: %s", i, path, exc)73    return samples74 75 76def resolve_dataset_path(dataset_dir: Path) -> Path:77    """Pick the best available training file.78 79    Preference order:80    1. train_annotated.jsonl  (LLM-augmented annotations)81    2. train.jsonl            (base annotations)82    """83    for name in ("train_annotated.jsonl", "train.jsonl", "final_train_v2.jsonl"):84        candidate = dataset_dir / name85        if candidate.is_file():86            logger.info("Using training dataset: %s", candidate)87            return candidate88    raise FileNotFoundError(89        f"No training dataset found in {dataset_dir}. "90        "Expected train_annotated.jsonl, train.jsonl, or final_train_v2.jsonl"91    )92 93 94def resolve_eval_path(dataset_dir: Path) -> Optional[Path]:95    """Find evaluation dataset if available."""96    for name in ("eval_annotated.jsonl", "eval.jsonl", "eval_notarial_only.jsonl"):97        candidate = dataset_dir / name98        if candidate.is_file():99            logger.info("Using eval dataset: %s", candidate)100            return candidate101    logger.info("No eval dataset found; training without evaluation.")102    return None103 104 105# ---------------------------------------------------------------------------106# Dataset statistics107# ---------------------------------------------------------------------------108 109def print_dataset_stats(samples: List[Dict[str, Any]], label: str = "Dataset") -> None:110    """Print summary statistics for a dataset split."""111    total_entities = 0112    label_counts: Counter = Counter()113    template_counts: Counter = Counter()114    text_lengths: List[int] = []115 116    for sample in samples:117        text = sample.get("text", sample.get("tokenized_text", ""))118        text_lengths.append(len(text))119 120        entities = sample.get("ner", sample.get("entities", []))121        total_entities += len(entities)122 123        for ent in entities:124            # Support both list format [start, end, label, text] and dict format125            if isinstance(ent, list) and len(ent) >= 3:126                label_counts[ent[2]] += 1127            elif isinstance(ent, dict):128                label_counts[ent.get("label", "UNKNOWN")] += 1129 130        template = sample.get("template", sample.get("template_id", "unknown"))131        template_counts[template] += 1132 133    print(f"\n{'=' * 60}")134    print(f"  {label} Statistics")135    print(f"{'=' * 60}")136    print(f"  Samples:          {len(samples):,}")137    print(f"  Total entities:   {total_entities:,}")138    if samples:139        print(f"  Avg entities/doc: {total_entities / len(samples):.1f}")140    if text_lengths:141        print(f"  Avg text length:  {sum(text_lengths) / len(text_lengths):,.0f} chars")142        print(f"  Min / Max length: {min(text_lengths):,} / {max(text_lengths):,} chars")143 144    if label_counts:145        print(f"\n  Entity labels ({len(label_counts)}):")146        for lbl, cnt in label_counts.most_common(30):147            print(f"    {lbl:<40s} {cnt:>6,}")148 149    if template_counts:150        print(f"\n  Templates ({len(template_counts)}):")151        for tpl, cnt in template_counts.most_common(20):152            print(f"    {tpl:<40s} {cnt:>6,}")153 154    print(f"{'=' * 60}\n")155 156 157# ---------------------------------------------------------------------------158# GLiNER training format conversion159# ---------------------------------------------------------------------------160 161def convert_to_gliner_format(samples: List[Dict[str, Any]]) -> List[Dict[str, Any]]:162    """Convert samples to GLiNER expected training format.163 164    GLiNER expects:165        {166            "tokenized_text": ["token1", "token2", ...],167            "ner": [[start_idx, end_idx, "LABEL"], ...]168        }169 170    If the input already has tokenized_text and ner in list format, pass through.171    Otherwise, convert from text + entities dict format.172    """173    converted = []174    skipped = 0175 176    for sample in samples:177        # Already in GLiNER format178        if "tokenized_text" in sample and "ner" in sample:179            converted.append({180                "tokenized_text": sample["tokenized_text"],181                "ner": sample["ner"],182            })183            continue184 185        # Convert from text + entities format186        text = sample.get("text", "")187        entities = sample.get("entities", sample.get("ner", []))188 189        if not text:190            skipped += 1191            continue192 193        # Simple whitespace tokenization (GLiNER handles subword internally)194        tokens = text.split()195        ner_annotations = []196 197        for ent in entities:198            if isinstance(ent, dict):199                start_char = ent.get("start", 0)200                end_char = ent.get("end", 0)201                label = ent.get("label", "UNKNOWN")202            elif isinstance(ent, list) and len(ent) >= 3:203                start_char, end_char, label = ent[0], ent[1], ent[2]204            else:205                continue206 207            # Map character offsets to token indices208            start_tok = _char_to_token(text, tokens, start_char)209            end_tok = _char_to_token(text, tokens, end_char - 1)210 211            if start_tok is not None and end_tok is not None:212                ner_annotations.append([start_tok, end_tok, label])213 214        converted.append({215            "tokenized_text": tokens,216            "ner": ner_annotations,217        })218 219    if skipped:220        logger.warning("Skipped %d samples with empty text.", skipped)221 222    logger.info(223        "Converted %d samples to GLiNER format (%d skipped).",224        len(converted), skipped,225    )226    return converted227 228 229def _char_to_token(text: str, tokens: List[str], char_offset: int) -> Optional[int]:230    """Map a character offset to a token index (whitespace tokenization)."""231    current_pos = 0232    for idx, token in enumerate(tokens):233        token_start = text.find(token, current_pos)234        if token_start == -1:235            token_start = current_pos236        token_end = token_start + len(token)237        if token_start <= char_offset < token_end:238            return idx239        current_pos = token_end240    return None241 242 243# ---------------------------------------------------------------------------244# Training245# ---------------------------------------------------------------------------246 247def run_training(248    train_samples: List[Dict[str, Any]],249    eval_samples: Optional[List[Dict[str, Any]]],250    output_dir: str,251    epochs: int,252    batch_size: int,253    lr: float,254    grad_accum: int,255    warmup_ratio: float,256    eval_steps: int,257    save_steps: int,258) -> None:259    """Run GLiNER fine-tuning."""260    from gliner import GLiNER261    from gliner.training import Trainer, TrainingArguments262 263    # Try to import the best data collator available264    data_collator = None265    try:266        from gliner.training import RelationExtractionSpanDataCollator267        data_collator = RelationExtractionSpanDataCollator268        logger.info("Using RelationExtractionSpanDataCollator.")269    except (ImportError, AttributeError):270        try:271            from gliner.training import DataCollatorWithPadding272            data_collator = DataCollatorWithPadding273            logger.info("Using DataCollatorWithPadding.")274        except (ImportError, AttributeError):275            logger.info("No specific data collator found; using Trainer default.")276 277    # Load base model278    logger.info("Loading base model: knowledgator/gliner-multitask-large-v0.5")279    model = GLiNER.from_pretrained("knowledgator/gliner-multitask-large-v0.5")280 281    # Convert datasets282    train_data = convert_to_gliner_format(train_samples)283    eval_data = convert_to_gliner_format(eval_samples) if eval_samples else None284 285    # Collect all entity labels from training data286    all_labels = set()287    for sample in train_data:288        for ent in sample.get("ner", []):289            if isinstance(ent, list) and len(ent) >= 3:290                all_labels.add(ent[2])291    logger.info("Unique entity labels in training data: %d", len(all_labels))292    logger.info("Labels: %s", sorted(all_labels))293 294    # Build training arguments (tuned for DGX Spark, 1 GPU)295    training_args = TrainingArguments(296        output_dir=output_dir,297        num_train_epochs=epochs,298        per_device_train_batch_size=batch_size,299        per_device_eval_batch_size=batch_size,300        gradient_accumulation_steps=grad_accum,301        learning_rate=lr,302        warmup_ratio=warmup_ratio,303        fp16=True,304        eval_steps=eval_steps if eval_data else None,305        save_steps=save_steps,306        save_total_limit=3,307        logging_steps=50,308        eval_strategy="steps" if eval_data else "no",309        save_strategy="steps",310        load_best_model_at_end=bool(eval_data),311        report_to="none",312        dataloader_num_workers=4,313        remove_unused_columns=False,314    )315 316    # Early stopping to prevent overfitting317    callbacks = []318    if eval_data:319        try:320            from transformers import EarlyStoppingCallback321            callbacks.append(EarlyStoppingCallback(early_stopping_patience=3))322            logger.info("Early stopping enabled (patience=3)")323        except ImportError:324            logger.warning("EarlyStoppingCallback not available")325 326    # Let the GLiNER Trainer handle collation via model.data_processor327    # Pass the data_processor's collate_fn which chains:328    #   collate_raw_batch → tokenize_inputs → create_labels329    model_collator = None330    if hasattr(model, 'data_processor'):331        dp = model.data_processor332        import torch as _torch333        def full_collate(batch):334            """Full GLiNER collation: raw samples → model-ready tensors."""335            raw = dp.collate_raw_batch(batch)336            tokenized = dp.tokenize_inputs(raw["tokens"], raw["classes_to_id"])337            labels = dp.create_labels(raw)338            tokenized["labels"] = labels339            # Add text_lengths (required by model.forward)340            tokenized["text_lengths"] = _torch.tensor(raw["seq_length"])341            # Add span info if available342            if raw.get("span_idx") is not None:343                tokenized["span_idx"] = raw["span_idx"]344                tokenized["span_mask"] = raw["span_mask"]345            return tokenized346        model_collator = full_collate347        logger.info("Using full GLiNER collation pipeline")348 349    # Build trainer kwargs350    trainer_kwargs: Dict[str, Any] = {351        "model": model,352        "args": training_args,353        "train_dataset": train_data,354        "callbacks": callbacks if callbacks else None,355    }356    if model_collator is not None:357        trainer_kwargs["data_collator"] = model_collator358    if eval_data:359        trainer_kwargs["eval_dataset"] = eval_data360 361    trainer = Trainer(**trainer_kwargs)362 363    # Train364    logger.info("Starting training: %d epochs, batch=%d, grad_accum=%d, lr=%s",365                epochs, batch_size, grad_accum, lr)366    trainer.train()367 368    # Save final model369    final_dir = os.path.join(output_dir, "final")370    os.makedirs(final_dir, exist_ok=True)371    model.save_pretrained(final_dir)372    logger.info("Model saved to: %s", final_dir)373 374    # Save label list375    labels_path = os.path.join(output_dir, "labels.json")376    with open(labels_path, "w", encoding="utf-8") as f:377        json.dump(sorted(all_labels), f, ensure_ascii=False, indent=2)378    logger.info("Labels saved to: %s", labels_path)379 380    print(f"\nTraining complete. Model saved to: {final_dir}")381    print(f"Labels saved to: {labels_path}")382 383 384# ---------------------------------------------------------------------------385# CLI386# ---------------------------------------------------------------------------387 388def parse_args() -> argparse.Namespace:389    parser = argparse.ArgumentParser(390        description="Fine-tune GLiNER on annotated notarial entity data.",391        formatter_class=argparse.RawDescriptionHelpFormatter,392        epilog="""393Examples:394  # Full training run395  python3 scripts/gliner_finetune.py \\396      --dataset-dir data/gliner \\397      --output-dir models/gliner-notarial-v1398 399  # Dry run: inspect dataset only400  python3 scripts/gliner_finetune.py \\401      --dataset-dir data/gliner \\402      --output-dir models/gliner-notarial-v1 \\403      --dry-run404 405  # Custom hyperparameters406  python3 scripts/gliner_finetune.py \\407      --dataset-dir data/gliner \\408      --output-dir models/gliner-notarial-v1 \\409      --epochs 10 --batch-size 4 --lr 5e-6410        """,411    )412    parser.add_argument(413        "--dataset-dir", type=str, required=True,414        help="Directory containing train_annotated.jsonl (or train.jsonl) "415             "and optionally eval_annotated.jsonl (or eval.jsonl).",416    )417    parser.add_argument(418        "--output-dir", type=str, required=True,419        help="Directory to save the fine-tuned model.",420    )421    parser.add_argument(422        "--epochs", type=int, default=5,423        help="Number of training epochs (default: 5).",424    )425    parser.add_argument(426        "--batch-size", type=int, default=8,427        help="Per-device training batch size (default: 8).",428    )429    parser.add_argument(430        "--lr", type=float, default=1e-5,431        help="Learning rate (default: 1e-5).",432    )433    parser.add_argument(434        "--grad-accum", type=int, default=4,435        help="Gradient accumulation steps (default: 4).",436    )437    parser.add_argument(438        "--warmup-ratio", type=float, default=0.1,439        help="Warmup ratio (default: 0.1).",440    )441    parser.add_argument(442        "--eval-steps", type=int, default=100,443        help="Evaluation interval in steps (default: 100).",444    )445    parser.add_argument(446        "--save-steps", type=int, default=200,447        help="Checkpoint save interval in steps (default: 200).",448    )449    parser.add_argument(450        "--dry-run", action="store_true",451        help="Load data and print statistics without training.",452    )453    return parser.parse_args()454 455 456def main() -> None:457    args = parse_args()458    dataset_dir = Path(args.dataset_dir)459 460    if not dataset_dir.is_dir():461        logger.error("Dataset directory not found: %s", dataset_dir)462        sys.exit(1)463 464    # Resolve dataset files465    try:466        train_path = resolve_dataset_path(dataset_dir)467    except FileNotFoundError as exc:468        logger.error(str(exc))469        sys.exit(1)470 471    eval_path = resolve_eval_path(dataset_dir)472 473    # Load datasets474    logger.info("Loading training data...")475    train_samples = load_jsonl(train_path)476    if not train_samples:477        logger.error("Training dataset is empty.")478        sys.exit(1)479 480    eval_samples = None481    if eval_path:482        logger.info("Loading evaluation data...")483        eval_samples = load_jsonl(eval_path)484        if not eval_samples:485            logger.warning("Eval dataset is empty; proceeding without evaluation.")486            eval_samples = None487 488    # Print stats489    print_dataset_stats(train_samples, label=f"Training ({train_path.name})")490    if eval_samples:491        print_dataset_stats(eval_samples, label=f"Evaluation ({eval_path.name})")492 493    # Dry run: stop here494    if args.dry_run:495        print("\n[DRY RUN] Dataset loaded and stats printed. No training performed.")496        print(f"  Train samples: {len(train_samples):,}")497        if eval_samples:498            print(f"  Eval samples:  {len(eval_samples):,}")499        print(f"  Output dir:    {args.output_dir}")500        print(f"  Epochs:        {args.epochs}")501        print(f"  Batch size:    {args.batch_size}")502        print(f"  Learning rate: {args.lr}")503        print(f"  Grad accum:    {args.grad_accum}")504        print(f"  Effective batch: {args.batch_size * args.grad_accum}")505        return506 507    # Check dependencies for actual training508    if not _check_gliner_available():509        logger.error(510            "GLiNER is not installed. Install with:\n"511            "  pip install gliner\n"512            "Or for training support:\n"513            "  pip install gliner[train]"514        )515        sys.exit(1)516 517    if not _check_training_available():518        logger.error(519            "GLiNER training module not available. Install training extras:\n"520            "  pip install gliner[train]\n"521            "Or install manually:\n"522            "  pip install gliner transformers[torch] datasets accelerate"523        )524        sys.exit(1)525 526    # Create output directory527    os.makedirs(args.output_dir, exist_ok=True)528 529    # Run training530    run_training(531        train_samples=train_samples,532        eval_samples=eval_samples,533        output_dir=args.output_dir,534        epochs=args.epochs,535        batch_size=args.batch_size,536        lr=args.lr,537        grad_accum=args.grad_accum,538        warmup_ratio=args.warmup_ratio,539        eval_steps=args.eval_steps,540        save_steps=args.save_steps,541    )542 543 544if __name__ == "__main__":545    main()546