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.
095
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 