COLE CI commited on
Commit
424c5d9
·
0 Parent(s):

deploy to HF Space

Browse files
This view is limited to 50 files because it contains too many changes.   See raw diff
Files changed (50) hide show
  1. .gitattributes +2 -0
  2. .gitignore +69 -0
  3. .idea/.gitignore +8 -0
  4. .pre-commit-config.yaml +27 -0
  5. CITATION.cff +40 -0
  6. CLAUDE.md +70 -0
  7. CODE_OF_CONDUCT.md +35 -0
  8. Dockerfile +55 -0
  9. LICENSE +21 -0
  10. Makefile +19 -0
  11. README.md +128 -0
  12. SECURITY.md +15 -0
  13. cole/__init__.py +6 -0
  14. cole/backend/__init__.py +0 -0
  15. cole/backend/evaluation.py +41 -0
  16. cole/backend/results/leaderboard.json +0 -0
  17. cole/backend/submission_api.py +251 -0
  18. cole/backend/submit_tools.py +37 -0
  19. cole/backend/validation_tools.py +90 -0
  20. cole/dataset/__init__.py +0 -0
  21. cole/dataset/dataset.py +101 -0
  22. cole/dataset/datasets_data.py +797 -0
  23. cole/dataset/prompt_builder.py +43 -0
  24. cole/docker_requirements.txt +18 -0
  25. cole/evaluation/__init__.py +0 -0
  26. cole/evaluation/evaluation_pipeline.py +159 -0
  27. cole/evaluation/evaluation_pipeline_private_llm.py +137 -0
  28. cole/evaluation/evaluation_pipeline_small.py +150 -0
  29. cole/evaluation/evaluation_pipeline_small_2.py +151 -0
  30. cole/evaluation/llm_evaluator.py +217 -0
  31. cole/evaluation/llm_factory.py +25 -0
  32. cole/evaluation/tools.py +33 -0
  33. cole/language_model/__init__.py +0 -0
  34. cole/language_model/anthropic_wrapper.py +20 -0
  35. cole/language_model/baseline.py +45 -0
  36. cole/language_model/cohere_wrapper.py +31 -0
  37. cole/language_model/deepseek_wrapper.py +30 -0
  38. cole/language_model/google_wrapper.py +22 -0
  39. cole/language_model/hugging_face_lm.py +161 -0
  40. cole/language_model/huggingface_language_model_factory.py +68 -0
  41. cole/language_model/init_function_calling.py +51 -0
  42. cole/language_model/language_model_abstraction.py +26 -0
  43. cole/language_model/mistral_wrapper.py +20 -0
  44. cole/language_model/open_ai_api_lm_wrapper.py +171 -0
  45. cole/language_model/open_ai_wrapper.py +32 -0
  46. cole/language_model/private_language_model_factory.py +117 -0
  47. cole/language_model/private_lm.py +64 -0
  48. cole/language_model/xai_wrapper.py +19 -0
  49. cole/metrics/__init__.py +0 -0
  50. cole/metrics/fquad_metric.py +108 -0
.gitattributes ADDED
@@ -0,0 +1,2 @@
 
 
 
1
+ *.jsonl filter=lfs diff=lfs merge=lfs -text
2
+ *.pdf filter=lfs diff=lfs merge=lfs -text
.gitignore ADDED
@@ -0,0 +1,69 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ .idea/COLE.iml
2
+ .idea/inspectionProfiles/profiles_settings.xml
3
+ .idea/misc.xml
4
+ .idea/modules.xml
5
+ .idea/vcs.xml
6
+ .idea/workspace.xml
7
+ .idea/*
8
+ /Benchmarks
9
+ /__pycache__
10
+ src/config.py
11
+ /Loaded_Models
12
+ /results
13
+ /hf_data
14
+
15
+
16
+ COLE_nlu_benchmark_site/node_modules
17
+ COLE_nlu_benchmark_site/.pnp
18
+ COLE_nlu_benchmark_site/.pnp.*
19
+ COLE_nlu_benchmark_site/.yarn/*
20
+ !.yarn/patches
21
+ !.yarn/plugins
22
+ !.yarn/releases
23
+ !.yarn/versions
24
+
25
+ # testing
26
+ /coverage
27
+ .coverage
28
+ coverage.xml
29
+ htmlcov/
30
+
31
+ # next.js
32
+ COLE_nlu_benchmark_site/.next/
33
+ COLE_nlu_benchmark_site/out/
34
+
35
+ # production
36
+ COLE_nlu_benchmark_site/build
37
+
38
+ # misc
39
+ .DS_Store
40
+ *.pem
41
+
42
+ # debug
43
+ npm-debug.log*
44
+ yarn-debug.log*
45
+ yarn-error.log*
46
+ .pnpm-debug.log*
47
+
48
+ # env files (can opt-in for committing if needed)
49
+ .env*
50
+
51
+ # vercel
52
+ .vercel
53
+
54
+ # typescript
55
+ *.tsbuildinfo
56
+ next-env.d.ts
57
+ /src/results
58
+ /src/__pycache__
59
+ /src/backend/__pycache__
60
+ /archives/Models/__pycache__
61
+ *.pyc
62
+ /Benchmarks_data
63
+ /src/light_eval_custom/results
64
+ /offline_evaluation/results
65
+ /frontend/node_modules
66
+ /.idea/*
67
+ /frontend/.next
68
+ /src/backend/results
69
+ /node_modules
.idea/.gitignore ADDED
@@ -0,0 +1,8 @@
 
 
 
 
 
 
 
 
 
1
+ # Default ignored files
2
+ /shelf/
3
+ /workspace.xml
4
+ # Editor-based HTTP Client requests
5
+ /httpRequests/
6
+ # Datasource local storage ignored files
7
+ /dataSources/
8
+ /dataSources.local.xml
.pre-commit-config.yaml ADDED
@@ -0,0 +1,27 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ repos:
2
+ - repo: https://github.com/psf/black
3
+ rev: 26.3.1
4
+ hooks:
5
+ - id: black
6
+
7
+ - repo: https://github.com/PyCQA/pylint
8
+ rev: v3.3.6
9
+ hooks:
10
+ - id: pylint
11
+ args: [cole/, tests/]
12
+ pass_filenames: false
13
+ additional_dependencies:
14
+ - fastapi
15
+ - uvicorn
16
+ - slowapi
17
+ - httpx
18
+ - pydantic
19
+
20
+ - repo: local
21
+ hooks:
22
+ - id: eslint
23
+ name: eslint
24
+ entry: bash -c 'cd frontend && npm run lint'
25
+ language: system
26
+ files: ^frontend/.*\.(js|jsx|ts|tsx)$
27
+ pass_filenames: false
CITATION.cff ADDED
@@ -0,0 +1,40 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ cff-version: 1.2.0
2
+ title: "COLE: a Comprehensive Benchmark for French Language Understanding Evaluation"
3
+ message: "If you use COLE in your research, please cite our paper."
4
+ type: software
5
+ authors:
6
+ - given-names: David
7
+ family-names: Beauchemin
8
+ affiliation: Université Laval
9
+ - given-names: Yan
10
+ family-names: Tremblay
11
+ affiliation: Université Laval
12
+ - given-names: Mohamed Amine
13
+ family-names: Youssef
14
+ affiliation: Université Laval
15
+ - given-names: Richard
16
+ family-names: Khoury
17
+ affiliation: Université Laval
18
+ repository-code: "https://github.com/GRAAL-Research/COLE"
19
+ url: "https://colebenchmark.org"
20
+ license: MIT
21
+ version: 1.0.0
22
+ date-released: "2025-10-07"
23
+ preferred-citation:
24
+ type: article
25
+ title: "COLE: a Comprehensive Benchmark for French Language Understanding Evaluation"
26
+ authors:
27
+ - given-names: David
28
+ family-names: Beauchemin
29
+ - given-names: Yan
30
+ family-names: Tremblay
31
+ - given-names: Mohamed Amine
32
+ family-names: Youssef
33
+ - given-names: Richard
34
+ family-names: Khoury
35
+ year: 2025
36
+ url: "https://arxiv.org/abs/2510.05046"
37
+ identifiers:
38
+ - type: other
39
+ value: "arXiv:2510.05046"
40
+ description: arXiv preprint
CLAUDE.md ADDED
@@ -0,0 +1,70 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # COLE — Quebec French NLU Benchmark
2
+
3
+ ## Project overview
4
+ COLE is a multidisciplinary Quebec French Natural Language Understanding benchmark with 23 tasks.
5
+ - **Backend**: FastAPI (port 8000) — submission API, HuggingFace dataset evaluation
6
+ - **Frontend**: Next.js 16 (port 8001) — leaderboard, submission form, bilingual (EN/FR)
7
+ - **Deployment**: Docker container on HuggingFace Spaces (nginx on port 7860 proxies both)
8
+
9
+ ## Quick start
10
+ ```bash
11
+ # Backend
12
+ export HF_TOKEN=hf_...
13
+ pip install -r cole/requirements.txt
14
+ uvicorn cole.backend.submission_api:app --host 0.0.0.0 --port 8000
15
+
16
+ # Frontend
17
+ cd frontend && npm ci && npm run dev
18
+
19
+ # Tests (requires HF_TOKEN for dataset access)
20
+ pip install -r tests/tests_requirements.txt
21
+ pytest
22
+ ```
23
+
24
+ ## Key directories
25
+ - `cole/backend/` — FastAPI app (`submission_api.py` is the entry point)
26
+ - `cole/dataset/` — HuggingFace dataset loading (repo: `graalul/COLE`)
27
+ - `cole/task/` — Task definitions and evaluation logic
28
+ - `cole/metrics/` — Metric wrappers (accuracy, F1, pearson, fquad); `math_metric.py` scores math answers by SymPy equivalence; `mixed_questions.py` evaluates JSONL files mixing several question types (Netquiz/COLE corpus)
29
+ - `frontend/src/app/` — Next.js App Router pages and components
30
+ - `tests/` — pytest tests (skip automatically without HF_TOKEN)
31
+
32
+ ## Commands
33
+ ```bash
34
+ make lint # Run pylint + eslint
35
+ make format # Check black formatting
36
+ make test # Run pytest
37
+ make build # Build frontend
38
+ make docker # Build Docker image
39
+ make all # lint + format + test + build
40
+ ```
41
+
42
+ ## CI/CD
43
+ - **Formatting**: `black --check .` (Python 3.12)
44
+ - **Linting**: `pylint cole/ tests/` (Python 3.10, 3.11, 3.12)
45
+ - **Tests**: `pytest` (Python 3.12, requires HF_TOKEN secret)
46
+ - **Frontend build**: `npm ci && npm run lint && npm run build`
47
+ - **Docker build**: builds and validates the Docker image
48
+ - **Deploy**: pushes to HuggingFace Spaces on main/dev push
49
+
50
+ ## Evaluating mixed-type question files
51
+ The Netquiz/COLE corpus stores several question types in a single JSONL file, each
52
+ row keeping its type in a `question_type` field. `cole.metrics.mixed_questions`
53
+ picks the metric per row (reusing the `fquad_metric` primitives) and reports a
54
+ per-type score plus two composites (unweighted, GLUE-style; and weighted by the
55
+ instance count per type):
56
+ ```bash
57
+ python -m cole.metrics.mixed_questions gold.jsonl
58
+ python -m cole.metrics.mixed_questions gold.jsonl --predictions preds.jsonl --output report.json
59
+ ```
60
+ Supported types and their metric: `single_choice`/`true_false` (accuracy),
61
+ `multiple_choice` (set F1), `short_answer` (FQuAD EM+F1, or SymPy-based math
62
+ equivalence via `math_metric.py` when the row's `subjects` flag it as
63
+ mathematical), `association` (pair F1), `categorization` (per-element accuracy),
64
+ `ordering` (exact match + pairwise concordance).
65
+
66
+ ## Conventions
67
+ - Python: black formatting (line-length 88), pylint score must be 10.0
68
+ - Frontend: ESLint with eslint-config-next flat config
69
+ - All components use `'use client'` directive (client-side i18n)
70
+ - Translations in `frontend/src/app/{en,fr}/translation.json`
CODE_OF_CONDUCT.md ADDED
@@ -0,0 +1,35 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Contributor Covenant Code of Conduct
2
+
3
+ ## Our Pledge
4
+
5
+ We as members, contributors, and leaders pledge to make participation in our
6
+ community a harassment-free experience for everyone, regardless of age, body
7
+ size, visible or invisible disability, ethnicity, sex characteristics, gender
8
+ identity and expression, level of experience, education, socio-economic status,
9
+ nationality, personal appearance, race, religion, or sexual identity
10
+ and orientation.
11
+
12
+ ## Our Standards
13
+
14
+ Examples of behavior that contributes to a positive environment:
15
+
16
+ * Using welcoming and inclusive language
17
+ * Being respectful of differing viewpoints and experiences
18
+ * Gracefully accepting constructive criticism
19
+ * Focusing on what is best for the community
20
+
21
+ Examples of unacceptable behavior:
22
+
23
+ * Trolling, insulting or derogatory comments, and personal or political attacks
24
+ * Public or private harassment
25
+ * Publishing others' private information without explicit permission
26
+ * Other conduct which could reasonably be considered inappropriate
27
+
28
+ ## Enforcement
29
+
30
+ Instances of abusive, harassing, or otherwise unacceptable behavior may be
31
+ reported to the project team at david.beauchemin@ift.ulaval.ca.
32
+
33
+ ## Attribution
34
+
35
+ This Code of Conduct is adapted from the [Contributor Covenant](https://www.contributor-covenant.org/), version 2.1.
Dockerfile ADDED
@@ -0,0 +1,55 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Stage 1: Build frontend
2
+ FROM node:20-slim AS frontend-build
3
+ WORKDIR /app/frontend
4
+ COPY frontend/package*.json ./
5
+ RUN npm ci
6
+ COPY frontend/ ./
7
+ RUN npm run build
8
+
9
+ # Stage 2: Final image with backend + built frontend
10
+ FROM python:3.12-slim
11
+
12
+ WORKDIR /app
13
+
14
+ # Install system dependencies (nginx, curl, Node.js runtime for Next.js)
15
+ RUN apt-get update && apt-get install -y --no-install-recommends \
16
+ nginx \
17
+ curl \
18
+ netcat-openbsd \
19
+ && curl -fsSL https://deb.nodesource.com/setup_20.x | bash - \
20
+ && apt-get install -y --no-install-recommends nodejs \
21
+ && rm -rf /var/lib/apt/lists/*
22
+
23
+ # Install Python dependencies
24
+ COPY cole/docker_requirements.txt /app/cole/
25
+ RUN pip install --no-cache-dir --upgrade pip wheel \
26
+ && pip install --no-cache-dir --prefer-binary pyarrow pandas numpy scipy fsspec aiohttp tqdm \
27
+ && pip install --no-cache-dir --prefer-binary -r /app/cole/docker_requirements.txt
28
+
29
+ # Copy backend source
30
+ COPY cole/ /app/cole/
31
+
32
+ # Copy built frontend from stage 1
33
+ COPY --from=frontend-build /app/frontend /app/frontend
34
+
35
+ # Copy config files
36
+ COPY nginx.conf /etc/nginx/nginx.conf
37
+ COPY start.sh /start.sh
38
+ RUN chmod +x /start.sh
39
+
40
+ # Create non-root user and set permissions
41
+ RUN useradd -m -u 1000 user \
42
+ && mkdir -p /app/.cache /var/lib/nginx /var/log/nginx /app/logs /run \
43
+ && touch /run/nginx.pid \
44
+ && chown -R user:user /app /var/lib/nginx /var/log/nginx /run
45
+
46
+ ENV HF_HOME=/app/.cache \
47
+ HF_DATASETS_CACHE=/app/.cache
48
+
49
+ USER user
50
+ EXPOSE 7860
51
+
52
+ HEALTHCHECK --interval=30s --timeout=10s --retries=3 \
53
+ CMD curl -f http://localhost:7860/ || exit 1
54
+
55
+ CMD ["sh", "/start.sh"]
LICENSE ADDED
@@ -0,0 +1,21 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ MIT License
2
+
3
+ Copyright (c) 2025 GRAAL Research Group, Université Laval
4
+
5
+ Permission is hereby granted, free of charge, to any person obtaining a copy
6
+ of this software and associated documentation files (the "Software"), to deal
7
+ in the Software without restriction, including without limitation the rights
8
+ to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
9
+ copies of the Software, and to permit persons to whom the Software is
10
+ furnished to do so, subject to the following conditions:
11
+
12
+ The above copyright notice and this permission notice shall be included in all
13
+ copies or substantial portions of the Software.
14
+
15
+ THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
16
+ IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
17
+ FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
18
+ AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
19
+ LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
20
+ OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
21
+ SOFTWARE.
Makefile ADDED
@@ -0,0 +1,19 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ .PHONY: lint format test build docker all
2
+
3
+ lint:
4
+ pylint cole/ tests/
5
+ cd frontend && npm run lint
6
+
7
+ format:
8
+ black --check .
9
+
10
+ test:
11
+ pytest
12
+
13
+ build:
14
+ cd frontend && npm ci && npm run build
15
+
16
+ docker:
17
+ docker build -t cole .
18
+
19
+ all: format lint test build
README.md ADDED
@@ -0,0 +1,128 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ title: COLE !
3
+ emoji: 🐳
4
+ colorFrom: purple
5
+ colorTo: gray
6
+ sdk: docker
7
+ app_port: 7860
8
+ ---
9
+
10
+ # COLE: Comprehensive Benchmark for Quebec French Language Understanding Evaluation
11
+
12
+ [![Website](https://img.shields.io/badge/Website-colebenchmark.org-blue)](https://colebenchmark.org/)
13
+ [![Paper](https://img.shields.io/badge/Paper-arXiv%3A2510.05046-b31b1b)](https://arxiv.org/abs/2510.05046)
14
+ [![Dataset](https://img.shields.io/badge/Dataset-HuggingFace-ffd21e)](https://huggingface.co/datasets/graalul/COLE-public)
15
+ [![Coverage](https://raw.githubusercontent.com/GRAAL-Research/COLE/badges/coverage-badge.svg)](https://github.com/GRAAL-Research/COLE/actions)
16
+
17
+ **COLE** is a comprehensive benchmark for evaluating Quebec French Natural Language Understanding (NLU). It includes 23 diverse tasks covering sentiment analysis, paraphrase detection, natural language inference, question answering, grammatical judgment, word sense disambiguation, and more — with a particular focus on linguistic phenomena relevant to the French language.
18
+
19
+ We benchmark 94 large language models (LLMs), providing an extensive analysis of the current state of Quebec French NLU. Our results highlight a significant performance gap between closed- and open-weight models and identify key challenging frontiers such as zero-shot extractive question answering, fine-grained word sense disambiguation, and understanding of regional language variations.
20
+
21
+ ## Links
22
+
23
+ - **Leaderboard**: [colebenchmark.org](https://colebenchmark.org/)
24
+ - **Paper**: [COLE: a Comprehensive Benchmark for Quebec French Language Understanding Evaluation (arXiv:2510.05046)](https://arxiv.org/abs/2510.05046)
25
+ - **Dataset**: [HuggingFace — graalul/COLE-public](https://huggingface.co/datasets/graalul/COLE-public)
26
+
27
+ ## Tasks
28
+
29
+ COLE consists of 23 tasks grouped by NLU capability:
30
+
31
+ ### Sentiment Analysis
32
+ | Task | Description | Test size |
33
+ |------|-------------|-----------|
34
+ | **Allocine** | Sentiment classification of French movie reviews (positive/negative) | 20,000 |
35
+ | **MMS-fr** | Sentiment analysis with 3 classes (positive, neutral, negative) | 63,190 |
36
+
37
+ ### Natural Language Inference (NLI)
38
+ | Task | Description | Test size |
39
+ |------|-------------|-----------|
40
+ | **DACCORD** | Semantic plausibility / contradiction detection of French sentences (binary) | 1,034 |
41
+ | **FraCaS** | NLI involving quantifiers, plurality, anaphora, and ellipsis | 346 |
42
+ | **GQNLI-fr** | NLI with quantifier logic (e.g., most, at least, more than half) | 30 |
43
+ | **LingNLI** | NLI corpus constructed with a linguist in the loop | 4,893 |
44
+ | **MNLI-nineeleven-Fr-MT** | French machine-translated MNLI using 9/11 context | 2,000 |
45
+ | **RTE3-Fr** | French version of RTE3 for textual entailment | 3,121 |
46
+ | **SICK-fr** | Sentence pair relatedness and entailment | 4,906 |
47
+ | **XNLI-fr** | Cross-lingual NLI in French | 5,010 |
48
+
49
+ ### Question Answering
50
+ | Task | Description | Test size |
51
+ |------|-------------|-----------|
52
+ | **FQuAD** | Extractive QA on high-quality French Wikipedia articles | 400 |
53
+ | **Fr-BoolQ** | Boolean question answering in French | 178 |
54
+ | **PIAF** | French extractive QA pairs | 384 |
55
+
56
+ ### Paraphrase Detection
57
+ | Task | Description | Test size |
58
+ |------|-------------|-----------|
59
+ | **PAWS-X** | Paraphrase identification from sentence pairs | 2,000 |
60
+ | **QFrBLiMP** | Semantic equivalence detection between sentence pairs | 2,290 |
61
+
62
+ ### Grammatical Judgment
63
+ | Task | Description | Test size |
64
+ |------|-------------|-----------|
65
+ | **MultiBLiMP-Fr** | Grammatical correctness from minimal pairs | 77 |
66
+ | **QFrCoLA** | Sentence acceptability in French (grammar, syntax) | 7,546 |
67
+
68
+ ### Semantic Similarity
69
+ | Task | Description | Test size |
70
+ |------|-------------|-----------|
71
+ | **STS22** | Document-level similarity of multilingual news articles | 72 |
72
+
73
+ ### Word Sense Disambiguation
74
+ | Task | Description | Test size |
75
+ |------|-------------|-----------|
76
+ | **WSD-Fr** | Disambiguating verb meanings in context | 3,121 |
77
+
78
+ ### Quebec French
79
+ | Task | Description | Test size |
80
+ |------|-------------|-----------|
81
+ | **QFrCoRE** | Matching Quebec French expressions to standard definitions | 4,633 |
82
+ | **QFrCoRT** | Matching Quebec French terms to standard definitions | 201 |
83
+
84
+ ### Coreference / Pronoun Resolution
85
+ | Task | Description | Test size |
86
+ |------|-------------|-----------|
87
+ | **Wino-X-LM** | Pronoun resolution with ambiguous referents | 2,793 |
88
+ | **Wino-X-MT** | Translation-based pronoun resolution with gendered pronouns | 2,988 |
89
+
90
+ ## Language
91
+
92
+ All data in COLE is in **French**.
93
+
94
+ ## Evaluating mixed-type question files
95
+
96
+ Some corpora (e.g. the Netquiz/COLE corpus) store several question types in a single
97
+ JSONL file, each row keeping its type in a `question_type` field. Since the right
98
+ metric differs from one question to the next, `cole.metrics.mixed_questions` selects
99
+ the metric per row and aggregates the results:
100
+
101
+ ```bash
102
+ # Predictions embedded in each row's "prediction" field
103
+ python -m cole.metrics.mixed_questions gold.jsonl
104
+
105
+ # Predictions in a separate file aligned line by line, JSON report written out
106
+ python -m cole.metrics.mixed_questions gold.jsonl --predictions preds.jsonl --output report.json
107
+ ```
108
+
109
+ It reuses COLE's metric primitives and reports a per-type score plus two composite
110
+ scores: an unweighted mean of the per-type scores (GLUE-style) and a mean weighted by
111
+ the number of instances of each type. Supported types: `single_choice`, `true_false`,
112
+ `multiple_choice`, `short_answer`, `association`, `categorization`, `ordering`.
113
+ Mathematical `short_answer` rows (flagged via their `subjects`) are scored with
114
+ SymPy-based equivalence so answers like `2^5` and `32` or `1/2` and `0.5` match.
115
+
116
+ ## Citation
117
+
118
+ If you use COLE in your research, please cite our paper:
119
+
120
+ ```bibtex
121
+ @article{beauchemin2025cole,
122
+ title={COLE: a Comprehensive Benchmark for Quebec French Language Understanding Evaluation},
123
+ author={Beauchemin, David and Tremblay, Yan and Youssef, Mohamed Amine and Khoury, Richard},
124
+ journal={arXiv preprint arXiv:2510.05046},
125
+ year={2025},
126
+ url={https://arxiv.org/abs/2510.05046}
127
+ }
128
+ ```
SECURITY.md ADDED
@@ -0,0 +1,15 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Security Policy
2
+
3
+ ## Supported Versions
4
+
5
+ | Version | Supported |
6
+ | ------- | ------------------ |
7
+ | 1.0.x | :white_check_mark: |
8
+
9
+ ## Reporting a Vulnerability
10
+
11
+ If you discover a security vulnerability, please report it responsibly by emailing **david.beauchemin@ift.ulaval.ca**.
12
+
13
+ Please do **not** open a public GitHub issue for security vulnerabilities.
14
+
15
+ We will acknowledge your report within 48 hours and provide a timeline for a fix.
cole/__init__.py ADDED
@@ -0,0 +1,6 @@
 
 
 
 
 
 
 
1
+ REPO_ID = "COLE-Graal/COLEGraal"
2
+ cole = "COLE-final"
3
+ boreal = "COLE-final-boreal"
4
+ complete = "COLE-finale-complete"
5
+ comparison = "Fr-comparison"
6
+ NA_VALUE = -1
cole/backend/__init__.py ADDED
File without changes
cole/backend/evaluation.py ADDED
@@ -0,0 +1,41 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import copy
2
+ import operator
3
+ from functools import reduce
4
+ from typing import List, Dict
5
+
6
+ from cole.task.task_factory import Task
7
+
8
+
9
+ def compute_tasks_ratings(tasks: List[Task], submission: Dict) -> Dict:
10
+ """
11
+ Method to compute the tasks ratings.
12
+ :param tasks: list of tasks
13
+ :param submission: submission dictionary
14
+ """
15
+
16
+ # We merge the tasks dictionary for simpler handling.
17
+ submission_copy = copy.deepcopy(submission)
18
+ submission_response = reduce(operator.ior, submission_copy.get("tasks"), {})
19
+
20
+ for task in tasks:
21
+ task_name = task.task_name
22
+
23
+ # We remove the prediction since we do not keep it in the response.
24
+ # Validation allows both "predictions" and "prediction"; pop whichever is present.
25
+ task_payload = submission_response.get(task_name)
26
+ if "predictions" in task_payload:
27
+ predictions = task_payload.pop("predictions")
28
+ else:
29
+ predictions = task_payload.pop("prediction")
30
+
31
+ ratings, warning = task.compute(predictions=predictions)
32
+ ratings.update({f"{task.metric_name}_warning": warning})
33
+ submission_response.get(task_name).update({f"{task.metric_name}": ratings})
34
+
35
+ # Final submission response where we unwrap the merge tasks dictionary into a list of dictionary.
36
+ submission_response = {
37
+ "model_name": submission.get("model_name"),
38
+ "model_url": submission.get("model_url"),
39
+ "tasks": [{key: value} for key, value in submission_response.items()],
40
+ }
41
+ return submission_response
cole/backend/results/leaderboard.json ADDED
The diff for this file is too large to render. See raw diff
 
cole/backend/submission_api.py ADDED
@@ -0,0 +1,251 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import glob
2
+ import json
3
+ import logging
4
+ import os
5
+ import uuid
6
+ from contextlib import asynccontextmanager
7
+ from datetime import datetime
8
+ from functools import lru_cache
9
+ from pathlib import Path
10
+ from typing import Dict, List, Any, Union
11
+
12
+ import huggingface_hub
13
+ from fastapi import FastAPI, UploadFile, Form, File, HTTPException
14
+ from fastapi.responses import JSONResponse
15
+ from fastapi.staticfiles import StaticFiles
16
+
17
+ from slowapi import Limiter
18
+ from slowapi.util import get_remote_address
19
+ from slowapi.errors import RateLimitExceeded
20
+ from starlette.middleware.cors import CORSMiddleware
21
+ from starlette.requests import Request
22
+
23
+ from cole.backend.evaluation import compute_tasks_ratings
24
+ from cole.backend.submit_tools import unzip_predictions_from_zip
25
+ from cole.dataset.datasets_data import preload_all_datasets
26
+ from cole.backend.validation_tools import (
27
+ validate_submission_tasks_name,
28
+ validate_submission_json,
29
+ validate_submission_template,
30
+ )
31
+ from cole.task.task import Task
32
+ from cole.task.task_factory import (
33
+ tasks_factory,
34
+ )
35
+
36
+ MAX_ZIP_SIZE_MB = 50
37
+
38
+ BASE_DIR = Path(__file__).resolve().parents[2]
39
+ RESULTS_DIR = BASE_DIR / "cole" / "backend" / "results"
40
+ RESULTS_DIR.mkdir(parents=True, exist_ok=True)
41
+ FRONTEND_DIR = BASE_DIR / "frontend"
42
+
43
+
44
+ @asynccontextmanager
45
+ async def lifespan(application: FastAPI = None): # pylint: disable=unused-argument
46
+ """Called before the backend comes online, is used to load datasets in memory."""
47
+ # Load the ML model
48
+ try:
49
+ token = os.environ.get("HF_TOKEN")
50
+ huggingface_hub.login(token=token)
51
+ preload_all_datasets()
52
+ except Exception as e:
53
+ error_message = f"The datasets could not be loaded : {e}"
54
+ logging.critical(error_message)
55
+
56
+ yield
57
+
58
+
59
+ limiter = Limiter(key_func=get_remote_address)
60
+ app = FastAPI(lifespan=lifespan)
61
+ app.state.limiter = limiter
62
+ app.add_exception_handler(
63
+ RateLimitExceeded,
64
+ lambda req, exc: JSONResponse(
65
+ status_code=429,
66
+ content={"detail": "Too many submissions. Please try again later."},
67
+ ),
68
+ )
69
+ app.mount("/results", StaticFiles(directory=str(RESULTS_DIR)), name="results")
70
+ front_end_info_message = f"The Front-end directory is: {FRONTEND_DIR}"
71
+ logging.info(front_end_info_message)
72
+
73
+ ALLOWED_ORIGINS = os.environ.get(
74
+ "CORS_ORIGINS",
75
+ "https://davebulaval-cole.hf.space,http://localhost:3000,http://localhost:8001",
76
+ ).split(",")
77
+
78
+ app.add_middleware(
79
+ CORSMiddleware,
80
+ allow_origins=ALLOWED_ORIGINS,
81
+ allow_methods=["GET", "POST"],
82
+ allow_headers=["*"],
83
+ )
84
+
85
+
86
+ @app.post("/submit")
87
+ @limiter.limit("5/minute")
88
+ async def submit(
89
+ request: Request, # pylint: disable=unused-argument # required by slowapi limiter
90
+ email: str = Form(...),
91
+ predictions_zip: UploadFile = File(...),
92
+ display_name: str = Form(...),
93
+ ):
94
+ """Route for making submissions with user generated results.
95
+ :param request : The incoming request (used for rate limiting)
96
+ :param email : The email of the user's submission
97
+ :param predictions_zip : The zip file of the user's predictions'
98
+ :param display_name : The display name associated with the user's submission'
99
+ """
100
+ logging.info("Starting submission")
101
+ if len(display_name) > 200:
102
+ raise HTTPException(
103
+ status_code=400, detail="Display name must be under 200 characters."
104
+ )
105
+ if len(email) > 320 or "@" not in email:
106
+ raise HTTPException(status_code=400, detail="Invalid email address.")
107
+ info_message = f"Submission from {email!r} as {display_name!r}."
108
+ logging.info(info_message)
109
+ zip_bytes = await predictions_zip.read()
110
+ if len(zip_bytes) > MAX_ZIP_SIZE_MB * 1024 * 1024:
111
+ raise HTTPException(
112
+ status_code=413, detail=f"ZIP file exceeds {MAX_ZIP_SIZE_MB}MB limit."
113
+ )
114
+ submission_json = unzip_predictions_from_zip(zip_bytes)
115
+
116
+ validate_submission_template(submission_json)
117
+ validate_submission_tasks_name(submission_json)
118
+ validate_submission_json(submission_json)
119
+
120
+ tasks: List[Task] = tasks_factory(submission_json)
121
+ logging.info("Computation started")
122
+ start = datetime.now()
123
+ submission_response = compute_tasks_ratings(tasks=tasks, submission=submission_json)
124
+ computation_time = datetime.now() - start
125
+ info_message = f"Computation ended in {computation_time}"
126
+ logging.info(info_message)
127
+ submission_id = str(uuid.uuid4())
128
+ submission_response.update(
129
+ {
130
+ "display_name": display_name,
131
+ "email": email,
132
+ "submission_id": submission_id,
133
+ }
134
+ )
135
+
136
+ out_path = RESULTS_DIR / f"{submission_id}.json"
137
+ with open(out_path, "w", encoding="utf-8") as f:
138
+ json.dump(submission_response, f, ensure_ascii=False, indent=2)
139
+
140
+ get_leaderboard_entries.cache_clear()
141
+
142
+ return JSONResponse(content=submission_response)
143
+
144
+
145
+ @lru_cache(maxsize=1)
146
+ def get_leaderboard_entries() -> List[Dict[str, Any]]:
147
+ """Returns all entries currently in the leaderboard.
148
+ Supporte aussi les fichiers JSON qui contiennent une LISTE d'entrées
149
+ et normalise les métriques 'plates' en groupes imbriqués pour le front.
150
+ """
151
+
152
+ def _wrap_flat_metrics(task_payload: Dict[str, Any]) -> Dict[str, Any]:
153
+ """
154
+ Si task_payload est 'plat' (ex: {"accuracy": 94.2}),
155
+ on le transforme en {"<group>": {...}} pour que le front puisse l'agréger.
156
+ Règles de nommage du groupe :
157
+ - présence de exact_match/f1 -> "fquad"
158
+ - sinon présence de acc/accuracy -> "accuracy"
159
+ - sinon présence de pearson/pearsonr/spearman -> "correlation"
160
+ - sinon -> "metrics"
161
+ Les valeurs >1 sont laissées telles quelles (le front normalise déjà % -> [0,1]).
162
+ """
163
+ if not isinstance(task_payload, dict):
164
+ return task_payload
165
+
166
+ # si c'est déjà "imbriqué" (une valeur est un dict), on ne touche pas
167
+ if any(isinstance(v, dict) for v in task_payload.values()):
168
+ return task_payload
169
+
170
+ keys = set(k.lower() for k in task_payload.keys())
171
+ if {"exact_match", "f1"} & keys:
172
+ group = "fquad"
173
+ elif {"accuracy", "acc"} & keys:
174
+ group = "accuracy"
175
+ elif {"pearson", "pearsonr", "spearman"} & keys:
176
+ group = "correlation"
177
+ else:
178
+ group = "metrics"
179
+
180
+ # Rien de spécial pour les warnings ici : le front les considère optionnels
181
+ # et s'attend à "<group>_warning" dans l'objet interne si on veut en fournir.
182
+ return {group: task_payload}
183
+
184
+ entries: List[Dict[str, Any]] = []
185
+
186
+ for filepath in glob.glob(str(RESULTS_DIR / "*.json")):
187
+ try:
188
+ with open(filepath, encoding="utf-8") as f:
189
+ data = json.load(f)
190
+
191
+ # Fonction interne qui traite UNE entrée (dict) au bon format minimal
192
+ def process_entry(entry: Dict[str, Any]) -> Union[Dict[str, Any], None]:
193
+ if not isinstance(entry, dict):
194
+ return None
195
+ if "model_name" not in entry or "tasks" not in entry:
196
+ return None
197
+
198
+ # Re-construire "results" comme le front s'y attend
199
+ results = {}
200
+ for task_obj in entry.get("tasks", []):
201
+ if not isinstance(task_obj, dict) or len(task_obj) != 1:
202
+ continue
203
+ task_name, payload = list(task_obj.items())[0]
204
+ normalized = _wrap_flat_metrics(payload)
205
+ results[task_name] = normalized
206
+
207
+ if not results:
208
+ return None
209
+
210
+ return {
211
+ "submission_id": entry.get("submission_id") or str(uuid.uuid4()),
212
+ "display_name": entry.get("display_name")
213
+ or entry.get("model_name")
214
+ or "Unnamed Model",
215
+ "email": entry.get("email", "N/A"),
216
+ "results": results,
217
+ }
218
+
219
+ # Le fichier peut contenir UNE entrée (dict) ou PLUSIEURS (list)
220
+ if isinstance(data, list):
221
+ for item in data:
222
+ processed = process_entry(item)
223
+ if processed:
224
+ entries.append(processed)
225
+ else:
226
+ processed = process_entry(data)
227
+ if processed:
228
+ entries.append(processed)
229
+
230
+ except Exception as e:
231
+ logging_message = f"Error processing file '{filepath}': {e}"
232
+ logging.error(logging_message)
233
+ continue
234
+
235
+ return entries
236
+
237
+
238
+ @app.get("/leaderboard")
239
+ async def leaderboard() -> List[Dict[str, Any]]:
240
+
241
+ return get_leaderboard_entries()
242
+
243
+
244
+ @app.get("/health")
245
+ async def health_check():
246
+ return {"status": "healthy", "message": "API is running."}
247
+
248
+
249
+ @app.get("/")
250
+ async def home():
251
+ return {"status": "working"}
cole/backend/submit_tools.py ADDED
@@ -0,0 +1,37 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import io
2
+ import json
3
+ import zipfile
4
+
5
+ from fastapi import HTTPException
6
+
7
+ MAX_DECOMPRESSED_SIZE_MB = 200
8
+
9
+
10
+ def unzip_predictions_from_zip(zip_bytes: bytes) -> dict:
11
+ """
12
+ Reads predictions.json directly from the ZIP in memory.
13
+ """
14
+ try:
15
+ zip_file = zipfile.ZipFile(io.BytesIO(zip_bytes))
16
+ except zipfile.BadZipFile as exc:
17
+ raise HTTPException(
18
+ 400, "The uploaded file is not a valid ZIP archive."
19
+ ) from exc
20
+
21
+ with zip_file as z:
22
+ if "predictions.json" not in z.namelist():
23
+ error_message = (
24
+ "The uploaded ZIP file does not contains a predictions.json file."
25
+ )
26
+ raise HTTPException(400, error_message)
27
+ info = z.getinfo("predictions.json")
28
+ if info.file_size > MAX_DECOMPRESSED_SIZE_MB * 1024 * 1024:
29
+ raise HTTPException(
30
+ 413,
31
+ f"Decompressed predictions.json exceeds {MAX_DECOMPRESSED_SIZE_MB}MB limit.",
32
+ )
33
+ with z.open("predictions.json") as f:
34
+ try:
35
+ return json.load(f)
36
+ except json.JSONDecodeError as exc:
37
+ raise HTTPException(400, "predictions.json is not valid JSON.") from exc
cole/backend/validation_tools.py ADDED
@@ -0,0 +1,90 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import logging
2
+ from typing import Dict, List
3
+
4
+ from fastapi import HTTPException
5
+
6
+ from cole.task.task_names import Tasks
7
+
8
+ tasks_name = [task.value for task in Tasks]
9
+
10
+
11
+ def validate_submission_template(dictionary: Dict) -> None:
12
+ """Ensures the dictionnary follows the correct format.
13
+ :param dictionary: Dictionary to validate."""
14
+ if dictionary.get("model_name", None) is None:
15
+ error = "The submission is missing a model name."
16
+ logging.error(error)
17
+ raise HTTPException(400, error)
18
+ if dictionary.get("model_url", None) is None:
19
+ error = "The submission is missing a model URL."
20
+ logging.error(error)
21
+ raise HTTPException(400, error)
22
+ if dictionary.get("tasks", None) is None:
23
+ error = "The submission is missing a tasks keyword."
24
+ logging.error(error)
25
+ raise HTTPException(400, error)
26
+
27
+ tasks = dictionary.get("tasks")
28
+ if not isinstance(tasks, List):
29
+ error = (
30
+ "The tasks keyword value must be a list of dictionaries where they key is the tasks "
31
+ "and value is a dictionary of predictions (in a list format). See our documentation for"
32
+ "a template."
33
+ )
34
+ logging.error(error)
35
+ raise HTTPException(400, error)
36
+
37
+ for task in tasks:
38
+ if not isinstance(task, dict) or len(task.keys()) != 1:
39
+ error = (
40
+ "Each task must be a dictionary of one element where the key is "
41
+ "the task name and the value is a list."
42
+ )
43
+ logging.error(error)
44
+ raise HTTPException(400, error)
45
+
46
+
47
+ def validate_submission_tasks_name(dictionary: Dict) -> None:
48
+ """
49
+ Validate if the submission JSON key are the tasks name.
50
+ """
51
+ for task in dictionary.get("tasks"):
52
+ key = list(task.keys())[0]
53
+ if key not in tasks_name:
54
+ error = f"Unknown key '{key}' in the submission JSON. The expected tasks are: {tasks_name}."
55
+ logging.error(error)
56
+ raise HTTPException(400, error)
57
+
58
+
59
+ def validate_submission_json(dictionary: Dict) -> None:
60
+ """Validates that the submitted json is in the correct format.
61
+ :param dictionary: Dictionary to validate."""
62
+ task_payload = dictionary.get("tasks")
63
+
64
+ for task in task_payload:
65
+ for task_name, payload in task.items():
66
+ if not isinstance(payload, dict):
67
+ error = (
68
+ "The tasks payload must be a dictionary in the format '{'prediction': [<predictions>]}' "
69
+ "for each task."
70
+ )
71
+ logging.error(error)
72
+ raise HTTPException(400, error)
73
+ if not ({"predictions", "prediction"} & payload.keys()):
74
+ # Empty payload (`{}`) used to slip through and crash compute_tasks_ratings
75
+ # with a KeyError -> HTTP 500. Reject it explicitly here.
76
+ error = f"The task '{task_name}' payload does not have the expected key: 'predictions'."
77
+ logging.error(error)
78
+ raise HTTPException(400, error)
79
+ for key, value in payload.items():
80
+ if key not in ["predictions", "prediction"]:
81
+ error = f"The task '{task_name}' payload does not have the expected key: 'predictions'."
82
+ logging.error(error)
83
+ raise HTTPException(400, error)
84
+ if not isinstance(value, list):
85
+ error = (
86
+ f"The task '{task_name}' predictions payload is not in a list format. "
87
+ r"The expected format is: '{'prediction': [<predictions>]}'"
88
+ )
89
+ logging.error(error)
90
+ raise HTTPException(400, error)
cole/dataset/__init__.py ADDED
File without changes
cole/dataset/dataset.py ADDED
@@ -0,0 +1,101 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from typing import Callable, Any, Union, List
2
+
3
+ from datasets import load_dataset
4
+
5
+
6
+ class Dataset:
7
+ """Class representing a usable dataset.
8
+ Allows dataset to be expressed as multiple forms, including as prompts, data or answers.
9
+ :param name : name of the dataset.
10
+ :param description : description of the dataset.
11
+ :param possible_ground_truths : the form that could be taken by ground truths.
12
+ :param hugging_face_repo : where to download the dataset on HuggingFace.
13
+ :param line_to_truth_fn : a function converting a dataset line to its truth value.
14
+ :param line_to_prompt_fn : a function converting a dataset line to a prompt for LLM inference.
15
+ :param line_to_data_fn : a function converting a dataset line to its data value for non LLM inference.
16
+ """
17
+
18
+ def __init__(
19
+ self,
20
+ name: str,
21
+ description: str,
22
+ possible_ground_truths: Union[List[str], List[int], List[float]],
23
+ hugging_face_repo: str,
24
+ line_to_truth_fn: Callable,
25
+ line_to_prompt_fn: Callable,
26
+ line_to_data_fn: Callable,
27
+ ):
28
+ self._dataset = None
29
+ self._ground_truths_cache = None
30
+ self.name = name
31
+ self.description = description
32
+ self.hugging_face_repo = hugging_face_repo
33
+ self.possible_ground_truths = possible_ground_truths
34
+ self.line_to_prompt_fn = line_to_prompt_fn
35
+ self.line_to_truth_fn = line_to_truth_fn
36
+ self.line_to_data_fn = line_to_data_fn
37
+
38
+ @property
39
+ def dataset(self):
40
+ self.load_data()
41
+ return self._dataset
42
+
43
+ def load_data(self):
44
+ if self._dataset is None:
45
+ self._dataset = load_dataset(
46
+ self.hugging_face_repo, name=self.name, split="test"
47
+ )
48
+
49
+ @property
50
+ def ground_truths(self) -> Union[List[str], List[int], List[float]]:
51
+ """The dataset's ground truths as a list (cached after first computation)"""
52
+ if self._ground_truths_cache is None:
53
+ self._ground_truths_cache = [
54
+ self.line_to_truth_fn(line) for line in self.dataset
55
+ ]
56
+ return self._ground_truths_cache
57
+
58
+ @property
59
+ def prompts(self) -> List[str]:
60
+ """The dataset's prompts as a list"""
61
+ return [self.line_to_prompt_fn(line) for line in self.dataset]
62
+
63
+ @property
64
+ def data(self) -> List[str]:
65
+ """The dataset's data as a list"""
66
+ return [self.line_to_data_fn(line) for line in self.dataset]
67
+
68
+ @property
69
+ def metadata(self) -> dict[str, Any]:
70
+ """The dataset's metadata as a dict"""
71
+ return {
72
+ "name": self.name,
73
+ "description": self.description,
74
+ "possible_ground_truths": str(self.possible_ground_truths),
75
+ "Prompt template": self.line_to_prompt_fn(self.EchoDict()),
76
+ }
77
+
78
+ @property
79
+ def metadata_string(self) -> str:
80
+ """The dataset's metadata as a string"""
81
+ lines = []
82
+ for key, value in self.metadata.items():
83
+ lines.append(f"{key}: {value}")
84
+ return "\n".join(lines)
85
+
86
+ def __len__(self):
87
+ return len(self.dataset)
88
+
89
+ def __getitem__(self, index: Union[int, slice]):
90
+ if isinstance(index, slice):
91
+ get_item_data = self.ground_truths[index.start : index.stop]
92
+ else:
93
+ get_item_data = self.ground_truths[index]
94
+
95
+ return get_item_data
96
+
97
+ class EchoDict:
98
+ """Helper class for building prompt templates,always returns the accessed key"""
99
+
100
+ def __getitem__(self, key):
101
+ return key
cole/dataset/datasets_data.py ADDED
@@ -0,0 +1,797 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from cole.dataset.dataset import Dataset
2
+ from cole.dataset.prompt_builder import PromptBuilder
3
+ from cole.task import COLE_REPOSITORY_NAME
4
+ from cole.task.task_names import COLETasks, BorealTasks
5
+
6
+ datasets = {
7
+ COLETasks.ALLOCINE.value: Dataset(
8
+ name=COLETasks.ALLOCINE.value,
9
+ description="Binary classification on sentiment analysis"
10
+ " of movie reviews, with reviews being either positive (1) or negative (0).",
11
+ possible_ground_truths=["0", "1"],
12
+ hugging_face_repo=COLE_REPOSITORY_NAME,
13
+ line_to_truth_fn=lambda line: line["label"],
14
+ line_to_prompt_fn=lambda line: PromptBuilder()
15
+ .add_premise("Cette phrase possède-t-elle un sentiment positif ou négatif ?")
16
+ .add_data(line["review"])
17
+ .add_end(
18
+ (
19
+ "Réponds "
20
+ "uniquement par 1 si la phrase est positive, réponds par 0 sinon. La réponse est :"
21
+ )
22
+ )
23
+ .build(),
24
+ line_to_data_fn=lambda line: line["review"],
25
+ ),
26
+ COLETasks.QFRCOLA.value: Dataset(
27
+ name=COLETasks.QFRCOLA.value,
28
+ description="Binary grammatical judgement : "
29
+ "Predicts whether a sentence is grammatically correct (1) or not. (0).",
30
+ possible_ground_truths=["0", "1"],
31
+ hugging_face_repo=COLE_REPOSITORY_NAME,
32
+ line_to_truth_fn=lambda line: line["label"],
33
+ line_to_prompt_fn=lambda line: PromptBuilder()
34
+ .add_premise("Juge si cette phrase est grammaticalement correcte :")
35
+ .add_data(line["sentence"])
36
+ .add_end(
37
+ (
38
+ "Réponds avec seulement 1 si la phrase est grammaticalement correcte, 0 sinon. La réponse est :"
39
+ )
40
+ )
41
+ .build(),
42
+ line_to_data_fn=lambda line: line["sentence"],
43
+ ),
44
+ COLETasks.QFRBLIMP.value: Dataset(
45
+ name=COLETasks.QFRBLIMP.value,
46
+ description="Choice task between two sentences : Choose the one which is grammatically correct.",
47
+ possible_ground_truths=["0", "1"],
48
+ hugging_face_repo=COLE_REPOSITORY_NAME,
49
+ line_to_truth_fn=lambda line: str(
50
+ line["label"]
51
+ ), # The label is return as a string.
52
+ line_to_prompt_fn=lambda line: (
53
+ PromptBuilder()
54
+ .add_premise("Laquelle de ces phrases est grammaticalement correcte ?")
55
+ .add_data(f"Phrase 0:{line['sentence_a']}")
56
+ .add_data(f"Phrase 1:{line['sentence_b']}")
57
+ .add_end(
58
+ "Réponds avec seulement 0 si la phrase 0 "
59
+ "est grammaticalement correcte, et uniquement 1 si la phrase 1 est grammaticalement "
60
+ "correcte. La réponse est :"
61
+ )
62
+ .build()
63
+ ),
64
+ line_to_data_fn=lambda line: {line["sentence_a"], line["sentence_b"]},
65
+ ),
66
+ COLETasks.GQNLI.value: Dataset(
67
+ name=COLETasks.GQNLI.value,
68
+ description="Natural language inference task : "
69
+ "predict the relation between two sentences (implication, neutral, contradiction).",
70
+ possible_ground_truths=["0", "1", "2"],
71
+ hugging_face_repo=COLE_REPOSITORY_NAME,
72
+ line_to_truth_fn=lambda line: line["label"],
73
+ line_to_prompt_fn=lambda line: PromptBuilder()
74
+ .add_premise(
75
+ "Quelle est la relation de la deuxième phrase par rapport à la première ?"
76
+ )
77
+ .add_data(line["premise"])
78
+ .add_data(line["hypothesis"])
79
+ .add_end(
80
+ (
81
+ "Réponds uniquement par :\n"
82
+ "0 - si la deuxième phrase implique la première,\n"
83
+ "1 - si la relation est neutre,\n"
84
+ "2 - s'il y a contradiction.\n"
85
+ "Réponds uniquement par 0, 1 ou 2. La réponse est :"
86
+ )
87
+ )
88
+ .build(),
89
+ line_to_data_fn=lambda line: {
90
+ "premise": line["premise"],
91
+ "hypothesis": line["hypothesis"],
92
+ },
93
+ ),
94
+ COLETasks.SICKFR.value: Dataset(
95
+ name=COLETasks.SICKFR.value,
96
+ description="Natural language inference task : "
97
+ "predict the relation between two sentences (implication, neutral, contradiction).",
98
+ possible_ground_truths=["0", "1", "2"],
99
+ hugging_face_repo=COLE_REPOSITORY_NAME,
100
+ line_to_truth_fn=lambda line: str(line["label"]),
101
+ line_to_prompt_fn=lambda line: PromptBuilder()
102
+ .add_premise("Détermine la relation entre les deux phrases suivantes :")
103
+ .add_data(f"Phrase A : {line['sentence_A']}\nPhrase B : {line['sentence_B']}")
104
+ .add_end(
105
+ "Réponds uniquement par 0, 1 ou 2 :\n"
106
+ "0 - si la deuxième phrase découle logiquement de la première,\n"
107
+ "1 - si leur relation est neutre,\n"
108
+ "2 - si les phrases se contredisent.\n"
109
+ "La réponse est :"
110
+ )
111
+ .build(),
112
+ line_to_data_fn=lambda line: {
113
+ "sentence_A": line["sentence_A"],
114
+ "sentence_B": line["sentence_B"],
115
+ },
116
+ ),
117
+ COLETasks.STS22.value: Dataset(
118
+ name=COLETasks.STS22.value,
119
+ description="Semantic textual similarity task : "
120
+ "Predict how similar two sentences are to each other (1 to 4).",
121
+ possible_ground_truths=["1", "2", "3", "4"],
122
+ hugging_face_repo=COLE_REPOSITORY_NAME,
123
+ line_to_truth_fn=lambda line: str(line["score"]),
124
+ line_to_prompt_fn=lambda line: PromptBuilder()
125
+ .add_premise(
126
+ "À quel point les deux phrases suivantes sont-elles similaires ? Donne une note entière de 1 à 4."
127
+ )
128
+ .add_data(f"Phrase 1 : {line['sentence1']}\nPhrase 2 : {line['sentence2']}")
129
+ .add_end(
130
+ "Réponds uniquement avec un nombre entier entre 1 (aucune similarité) et 4 (équivalence parfaite). "
131
+ "La réponse est :"
132
+ )
133
+ .build(),
134
+ line_to_data_fn=lambda line: {
135
+ "sentence1": line["sentence1"],
136
+ "sentence2": line["sentence2"],
137
+ },
138
+ ),
139
+ COLETasks.PAWS_X.value: Dataset(
140
+ name=COLETasks.PAWS_X.value,
141
+ description="Binary classification task : "
142
+ "Predict if two sentences have the same meaning (1) or not (0).",
143
+ possible_ground_truths=["0", "1"],
144
+ hugging_face_repo=COLE_REPOSITORY_NAME,
145
+ line_to_truth_fn=lambda line: line["label"],
146
+ line_to_prompt_fn=lambda line: PromptBuilder()
147
+ .add_premise(
148
+ "Les deux phrases suivantes veulent-elles dire la même chose, ou ont-elles des significations différentes ?"
149
+ )
150
+ .add_data(line["sentence1"])
151
+ .add_data(line["sentence2"])
152
+ .add_end(
153
+ (
154
+ "Réponds seulement 1 si les deux phrases ont la même signification, 0 sinon. La réponse est :"
155
+ )
156
+ )
157
+ .build(),
158
+ line_to_data_fn=lambda line: {
159
+ "sentence1": line["sentence1"],
160
+ "sentence2": line["sentence2"],
161
+ },
162
+ ),
163
+ COLETasks.PIAF.value: Dataset(
164
+ name=COLETasks.PIAF.value,
165
+ description="Extractive question answering task : Extract a question's answer from a given context.",
166
+ possible_ground_truths=[],
167
+ hugging_face_repo=COLE_REPOSITORY_NAME,
168
+ line_to_truth_fn=lambda line: line["answers"],
169
+ line_to_prompt_fn=lambda line: PromptBuilder()
170
+ .add_premise(
171
+ "Tu vas recevoir un contexte suivi d'une question.\n"
172
+ "Ta tâche est d'extraire **mot pour mot** le passage du contexte qui répond le mieux à la question.\n"
173
+ "N'invente rien. Ne reformule pas.\n"
174
+ "Réponds **en copiant uniquement** un extrait exact du texte ci-dessus."
175
+ )
176
+ .add_data(f"Contexte : {line['context']}")
177
+ .add_data(f"Question : {line['question']}")
178
+ .add_end(
179
+ "Réponds uniquement par un passage extrait du contexte. La réponse est :"
180
+ )
181
+ .build(),
182
+ line_to_data_fn=lambda line: {
183
+ "context": line["context"],
184
+ "question": line["question"],
185
+ },
186
+ ),
187
+ COLETasks.FQUAD.value: Dataset(
188
+ name=COLETasks.FQUAD.value,
189
+ description="Extractive question answering task : Extract a question's answer from a given context.",
190
+ possible_ground_truths=[],
191
+ hugging_face_repo=COLE_REPOSITORY_NAME,
192
+ line_to_truth_fn=lambda line: line["answers"],
193
+ line_to_prompt_fn=lambda line: PromptBuilder()
194
+ .add_premise(
195
+ "Tu vas recevoir un contexte suivi d'une question.\n"
196
+ "Ta tâche est d'extraire **mot pour mot** le passage du contexte qui répond le mieux à la question.\n"
197
+ "N'invente rien. Ne reformule pas.\n"
198
+ "Réponds **en copiant uniquement** un extrait exact du texte ci-dessus."
199
+ )
200
+ .add_data(f"Contexte : {line['context']}")
201
+ .add_data(f"Question : {line['question']}")
202
+ .add_end(
203
+ "Réponds uniquement par un passage extrait du contexte. La réponse est :"
204
+ )
205
+ .build(),
206
+ line_to_data_fn=lambda line: {
207
+ "context": line["context"],
208
+ "question": line["question"],
209
+ },
210
+ ),
211
+ COLETasks.XNLI.value: Dataset(
212
+ name=COLETasks.XNLI.value,
213
+ description="Natural language inference task : "
214
+ "predict the relation between two sentences (implication, neutral, contradiction).",
215
+ possible_ground_truths=["0", "1", "2"],
216
+ hugging_face_repo=COLE_REPOSITORY_NAME,
217
+ line_to_truth_fn=lambda line: str(line["label"]),
218
+ line_to_prompt_fn=lambda line: PromptBuilder()
219
+ .add_premise(
220
+ "Quelle est la relation de la deuxième phrase par rapport à la première ?"
221
+ )
222
+ .add_data(rf"premise : {line['premise']}\n" f"sentence 2: {line['hypothesis']}")
223
+ .add_end(
224
+ (
225
+ "Réponds uniquement par :\n"
226
+ "0 - si la deuxième phrase implique la première,\n"
227
+ "1 - si la relation est neutre,\n"
228
+ "2 - s'il y a contradiction.\n"
229
+ "Réponds uniquement par 0, 1 ou 2. La réponse est :"
230
+ )
231
+ )
232
+ .build(),
233
+ line_to_data_fn=lambda line: {
234
+ "premise": line["premise"],
235
+ "hypothesis": line["hypothesis"],
236
+ },
237
+ ),
238
+ COLETasks.QFRCORE.value: Dataset(
239
+ name=COLETasks.QFRCORE.value,
240
+ description="Definition matching task : "
241
+ "Match the Quebec expression with its definition from a list.",
242
+ possible_ground_truths=[str(i) for i in range(10)],
243
+ hugging_face_repo=COLE_REPOSITORY_NAME,
244
+ line_to_truth_fn=lambda line: str(line["correct_index"]),
245
+ line_to_prompt_fn=lambda line: PromptBuilder()
246
+ .add_premise(
247
+ f"Que veut dire cette expression québécoise « {line['expression']} » ?"
248
+ )
249
+ .add_data(
250
+ "\n".join(
251
+ f"{idx} - {definition}"
252
+ for idx, definition in enumerate(line["choices"])
253
+ )
254
+ )
255
+ .add_end(
256
+ (
257
+ "Réponds uniquement par l'index, débutant à zéro, "
258
+ "de la bonne définition parmi la liste ci-dessus. Par exemple, si la "
259
+ "troisième phrase correspond à l'expression, la réponse sera 2. La réponse est :"
260
+ )
261
+ )
262
+ .build(),
263
+ line_to_data_fn=lambda line: {
264
+ "expression": line["expression"],
265
+ "choices": line["choices"],
266
+ },
267
+ ),
268
+ COLETasks.QFRCORT.value: Dataset(
269
+ name=COLETasks.QFRCORT.value,
270
+ description="Definition matching task : "
271
+ "Match the Quebec term with its definition from a list.",
272
+ possible_ground_truths=[str(i) for i in range(10)],
273
+ hugging_face_repo=COLE_REPOSITORY_NAME,
274
+ line_to_truth_fn=lambda line: str(line["correct_index"]),
275
+ line_to_prompt_fn=lambda line: PromptBuilder()
276
+ .add_premise(
277
+ f"Qu'est-ce que ça veut dire ce terme québécois « {line['terme']} » ?"
278
+ )
279
+ .add_data(
280
+ "\n".join(
281
+ f"{idx} - {definition}"
282
+ for idx, definition in enumerate(line["choices"])
283
+ )
284
+ )
285
+ .add_end(
286
+ (
287
+ "Réponds uniquement par l'index, débutant à zéro, "
288
+ "de la bonne définition parmi la liste ci-dessus. La réponse est :"
289
+ )
290
+ )
291
+ .build(),
292
+ line_to_data_fn=lambda line: {
293
+ "terme": line["terme"],
294
+ "choices": line["choices"],
295
+ },
296
+ ),
297
+ COLETasks.FRCOE.value: Dataset(
298
+ name=COLETasks.FRCOE.value,
299
+ description="Definition matching task : "
300
+ "Match the French expression with its definition from a list.",
301
+ possible_ground_truths=[str(i) for i in range(10)],
302
+ hugging_face_repo=COLE_REPOSITORY_NAME,
303
+ line_to_truth_fn=lambda line: str(line["correct_index"]),
304
+ line_to_prompt_fn=lambda line: PromptBuilder()
305
+ .add_premise(
306
+ f"Que veut dire cette expression française « {line['expression']} » ?"
307
+ )
308
+ .add_data(
309
+ "\n".join(
310
+ f"{idx} - {definition}"
311
+ for idx, definition in enumerate(line["choices"])
312
+ )
313
+ )
314
+ .add_end(
315
+ (
316
+ "Réponds uniquement par l'index, débutant à zéro, "
317
+ "de la bonne définition parmi la liste ci-dessus. Par exemple, si la "
318
+ "troisième phrase correspond à l'expression, la réponse sera 2. La réponse est :"
319
+ )
320
+ )
321
+ .build(),
322
+ line_to_data_fn=lambda line: {
323
+ "expression": line["expression"],
324
+ "choices": line["choices"],
325
+ },
326
+ ),
327
+ COLETasks.DACCORD.value: Dataset(
328
+ name=COLETasks.DACCORD.value,
329
+ description="Paraphrase detection task :"
330
+ "Predict whether the two sentences are compatible (0) "
331
+ "or contradict each other (1).",
332
+ possible_ground_truths=["0", "1"],
333
+ hugging_face_repo=COLE_REPOSITORY_NAME,
334
+ line_to_truth_fn=lambda line: str(line["label"]),
335
+ line_to_prompt_fn=lambda line: (
336
+ PromptBuilder()
337
+ .add_premise("Détermine la relation entre les deux phrases suivantes :")
338
+ .add_data(f"Première phrase : {line['premise']}")
339
+ .add_data(f"Deuxième phrase : {line['hypothesis']}")
340
+ .add_end(
341
+ "Réponds uniquement par :\n"
342
+ "0 - si les deux phrases sont compatibles (elles expriment la même information ou sont cohérentes),\n"
343
+ "1 - s'il y a contradiction entre les deux phrases.\n"
344
+ "Réponds uniquement par 0 ou 1. La réponse est :"
345
+ )
346
+ .build()
347
+ ),
348
+ line_to_data_fn=lambda line: {
349
+ "premise": line["premise"],
350
+ "hypothesis": line["hypothesis"],
351
+ },
352
+ ),
353
+ COLETasks.FRENCH_BOOLQ.value: Dataset(
354
+ name=COLETasks.FRENCH_BOOLQ.value,
355
+ description="Binary question answering task : "
356
+ "Answer whether the context allows answering 'yes' to the question (1)"
357
+ "or, if the context only allows answering 'no' "
358
+ "to the question or does not answer the question. (0).",
359
+ possible_ground_truths=["0", "1"],
360
+ hugging_face_repo=COLE_REPOSITORY_NAME,
361
+ line_to_truth_fn=lambda line: str(line["label"]),
362
+ line_to_prompt_fn=lambda line: (
363
+ PromptBuilder()
364
+ .add_premise(
365
+ "Lis le passage suivant et réponds à la question en te basant uniquement sur le texte :\n"
366
+ "- Si le passage permet d'affirmer que la réponse à la question est oui, réponds 1.\n"
367
+ "- Sinon, si la réponse est non ou que le passage ne permet pas de répondre à la question, réponds 0."
368
+ )
369
+ .add_data(f"Passage : {line['passage']}")
370
+ .add_data(f"Question : {line['question']}")
371
+ .add_end("La réponse est :")
372
+ .build()
373
+ ),
374
+ line_to_data_fn=lambda line: {
375
+ "question": line["question"],
376
+ "passage": line["passage"],
377
+ },
378
+ ),
379
+ COLETasks.MNLI_NINEELEVEN_FR_MT.value: Dataset(
380
+ name=COLETasks.MNLI_NINEELEVEN_FR_MT.value,
381
+ description="Natural language inference task : "
382
+ "predict the relation between two sentences (implication, neutral, contradiction).",
383
+ possible_ground_truths=["0", "1", "2"],
384
+ hugging_face_repo=COLE_REPOSITORY_NAME,
385
+ line_to_truth_fn=lambda line: line["label"],
386
+ line_to_prompt_fn=lambda line: PromptBuilder()
387
+ .add_premise(
388
+ "Quelle est la relation de la deuxième phrase par rapport à la première ?"
389
+ )
390
+ .add_data(line["premise"])
391
+ .add_data(line["hypothesis"])
392
+ .add_end(
393
+ (
394
+ "Réponds uniquement par :\n"
395
+ "0 - si la deuxième phrase implique la première,\n"
396
+ "1 - si la relation est neutre,\n"
397
+ "2 - s'il y a contradiction.\n"
398
+ "Réponds uniquement par 0, 1 ou 2. La réponse est :"
399
+ )
400
+ )
401
+ .build(),
402
+ line_to_data_fn=lambda line: {
403
+ "premise": line["premise"],
404
+ "hypothesis": line["hypothesis"],
405
+ },
406
+ ),
407
+ COLETasks.RTE3_FRENCH.value: Dataset(
408
+ name=COLETasks.RTE3_FRENCH.value,
409
+ description="Natural language inference task : "
410
+ "predict the relation between two sentences (entailment, neutral, contradiction)",
411
+ possible_ground_truths=["0", "1", "2"],
412
+ hugging_face_repo=COLE_REPOSITORY_NAME,
413
+ line_to_truth_fn=lambda line: str(line["label"]),
414
+ line_to_prompt_fn=lambda line: (
415
+ PromptBuilder()
416
+ .add_premise(
417
+ "Lis le texte suivant et détermine la relation de l'énoncé par rapport au texte."
418
+ )
419
+ .add_data(f"Texte : {line['premise']}")
420
+ .add_data(f"Énoncé : {line['hypothesis']}")
421
+ .add_end(
422
+ "Réponds uniquement par 0, 1 ou 2 :\n"
423
+ "0 - si l'énoncé découle logiquement du texte (entailment),\n"
424
+ "1 - si la relation est neutre,\n"
425
+ "2 - s'il y a contradiction.\n"
426
+ "La réponse est :"
427
+ )
428
+ .build()
429
+ ),
430
+ line_to_data_fn=lambda line: {
431
+ "premise": line["premise"],
432
+ "hypothesis": line["hypothesis"],
433
+ },
434
+ ),
435
+ COLETasks.WINO_X_LM.value: Dataset(
436
+ name=COLETasks.WINO_X_LM.value,
437
+ description=(
438
+ "Pronoun resolution task : predict the correct referent (1 or 2) "
439
+ "of a pronoun in a sentence by choosing between two candidates."
440
+ ),
441
+ possible_ground_truths=["1", "2"],
442
+ hugging_face_repo=COLE_REPOSITORY_NAME,
443
+ line_to_truth_fn=lambda line: str(line["answer"]),
444
+ line_to_prompt_fn=lambda line: PromptBuilder()
445
+ .add_premise(
446
+ 'Voici une phrase en anglais contenant le pronom "it" dans un sens ambigu et sa traduction en français.'
447
+ )
448
+ .add_data(f"Phrase (originale en anglais) : {line['sentence']}")
449
+ .add_data(
450
+ f"Traduction en français (le pronom est caché par '_' ) : {line['context_fr']}"
451
+ )
452
+ .add_data("À quoi renvoie ce pronom ? Voici les choix: ")
453
+ .add_data(f"1 : {line['option1_fr']}")
454
+ .add_data(f"2 : {line['option2_fr']}")
455
+ .add_end("Réponds uniquement par 1 ou 2. La réponse est :")
456
+ .build(),
457
+ line_to_data_fn=lambda line: {
458
+ "sentence": line["sentence"],
459
+ "translation": line["context_fr"],
460
+ "referent1": line["option1_fr"],
461
+ "referent2": line["option2_fr"],
462
+ },
463
+ ),
464
+ COLETasks.WINO_X_MT.value: Dataset(
465
+ name="wino_x_mt",
466
+ description=(
467
+ "Pronoun resolution based on translations: choose between two French translations of an English "
468
+ "sentence with an ambiguous pronoun. The goal is to identify which of the two translations uses "
469
+ "the correct pronoun (he or she) based on the correct referent."
470
+ ),
471
+ possible_ground_truths=["1", "2"],
472
+ hugging_face_repo=COLE_REPOSITORY_NAME,
473
+ line_to_truth_fn=lambda line: str(line["answer"]),
474
+ line_to_prompt_fn=lambda line: PromptBuilder()
475
+ .add_premise(
476
+ "Voici deux traductions d’une phrase anglaise contenant un pronom ambigu :"
477
+ )
478
+ .add_data(f"Phrase originale : {line['sentence']}")
479
+ .add_data(f"Traduction 1 (avec '{line['pronoun1']}') : {line['translation1']}")
480
+ .add_data(f"Traduction 2 (avec '{line['pronoun2']}') : {line['translation2']}")
481
+ .add_end(
482
+ "Quelle traduction utilise le bon pronom en fonction du référent visé dans la phrase originale ?\n"
483
+ "Réponds uniquement par 1 si la traduction 1 est correcte, ou 2 si la traduction 2 est correcte.\n"
484
+ "La réponse est :"
485
+ )
486
+ .build(),
487
+ line_to_data_fn=lambda line: {
488
+ "sentence": line["sentence"],
489
+ "translation1": line["translation1"],
490
+ "translation2": line["translation2"],
491
+ "pronoun1": line["pronoun1"],
492
+ "pronoun2": line["pronoun2"],
493
+ },
494
+ ),
495
+ COLETasks.MULTIBLIMP.value: Dataset(
496
+ name=COLETasks.MULTIBLIMP.value,
497
+ description="Choice task between two sentences : Choose the one which is grammatically correct.",
498
+ possible_ground_truths=["0", "1"],
499
+ hugging_face_repo=COLE_REPOSITORY_NAME,
500
+ line_to_truth_fn=lambda line: str(
501
+ line["label"]
502
+ ), # The label is return as a string.
503
+ line_to_prompt_fn=lambda line: (
504
+ PromptBuilder()
505
+ .add_premise("Laquelle de ces phrases est grammaticalement correcte ?")
506
+ .add_data(f"Phrase 0:{line['sentence_a']}")
507
+ .add_data(f"Phrase 1:{line['sentence_b']}")
508
+ .add_end(
509
+ "Réponds avec seulement 0 si la phrase 0 "
510
+ "est grammaticalement correcte, et uniquement 1 si la phrase 1 est grammaticalement "
511
+ "correcte. La réponse est :"
512
+ )
513
+ .build()
514
+ ),
515
+ line_to_data_fn=lambda line: {line["sentence_a"], line["sentence_b"]},
516
+ ),
517
+ COLETasks.FRACAS.value: Dataset(
518
+ name=COLETasks.FRACAS.value,
519
+ description="Natural language inference task : "
520
+ "predict the relation between two sentences (implication, neutral, contradiction).",
521
+ possible_ground_truths=["0", "1", "2"],
522
+ hugging_face_repo=COLE_REPOSITORY_NAME,
523
+ line_to_truth_fn=lambda line: line["label"],
524
+ line_to_prompt_fn=lambda line: PromptBuilder()
525
+ .add_premise(
526
+ "Quelle est la relation de la deuxième phrase par rapport à la première ?"
527
+ )
528
+ .add_data(line["premise"])
529
+ .add_data(line["hypothesis"])
530
+ .add_end(
531
+ (
532
+ "Réponds uniquement par :\n"
533
+ "0 - si la deuxième phrase implique la première,\n"
534
+ "1 - si la relation est neutre,\n"
535
+ "2 - s'il y a contradiction.\n"
536
+ "Réponds uniquement par 0, 1 ou 2. La réponse est :"
537
+ )
538
+ )
539
+ .build(),
540
+ line_to_data_fn=lambda line: {
541
+ "premise": line["premise"],
542
+ "hypothesis": line["hypothesis"],
543
+ },
544
+ ),
545
+ COLETasks.MMS.value: Dataset(
546
+ name=COLETasks.MMS.value,
547
+ description="A sentiment analysis task for classifying text as positive (2), negative (0), or neutral (1).",
548
+ possible_ground_truths=["0", "1", "2"],
549
+ hugging_face_repo=COLE_REPOSITORY_NAME,
550
+ line_to_truth_fn=lambda line: line["label"],
551
+ line_to_prompt_fn=lambda line: PromptBuilder()
552
+ .add_premise("Quel est le sentiment de cette phrase?")
553
+ .add_data(line["text"])
554
+ .add_end(
555
+ (
556
+ "Réponds uniquement par :\n"
557
+ "0 - si la phrase est négative,\n"
558
+ "1 - si la phrase est neutre,\n"
559
+ "2 - si la phrase est positive.\n"
560
+ "Réponds uniquement par 0, 1 ou 2. La réponse est :"
561
+ )
562
+ )
563
+ .build(),
564
+ line_to_data_fn=lambda line: {
565
+ "text": line["text"],
566
+ },
567
+ ),
568
+ COLETasks.WSD.value: Dataset(
569
+ name=COLETasks.WSD.value,
570
+ description="Extractive word sense disambiguation : Extract an ambiguous word in a sentence.",
571
+ possible_ground_truths=[],
572
+ hugging_face_repo=COLE_REPOSITORY_NAME,
573
+ line_to_truth_fn=lambda line: line["label"],
574
+ line_to_prompt_fn=lambda line: PromptBuilder()
575
+ .add_premise(
576
+ "Tu vas recevoir une phrase contenant un mot ambigu ainsi que les étiquettes du 'part-of-speech tagging "
577
+ "(PoS)' pour chaque mot de la phrase. Le mot ambigu peut être un verbe ou un adjectif.\n"
578
+ "Ta tâche est d’indiquer **exactement** ce mot ambigu dans la phrase, sans rien ajouter ni reformuler.\n"
579
+ "Réponds uniquement avec le mot ambigu identifié."
580
+ )
581
+ .add_data(f"Phrase : {line['sentence']}")
582
+ .add_data(f"Part-of-speech tagging: {line['pos_tag_labels']}")
583
+ .add_end("La réponse est :")
584
+ .build(),
585
+ line_to_data_fn=lambda line: {
586
+ "sentence": line["sentence"],
587
+ "pos_tag_labels": line["pos_tag_labels"],
588
+ },
589
+ ),
590
+ COLETasks.LINGNLI.value: Dataset(
591
+ name=COLETasks.LINGNLI.value,
592
+ description="Natural language inference task : "
593
+ "predict the relation between two sentences (implication, neutral, contradiction).",
594
+ possible_ground_truths=["0", "1", "2"],
595
+ hugging_face_repo=COLE_REPOSITORY_NAME,
596
+ line_to_truth_fn=lambda line: line["label"],
597
+ line_to_prompt_fn=lambda line: PromptBuilder()
598
+ .add_premise(
599
+ "Quelle est la relation de la deuxième phrase par rapport à la première ?"
600
+ )
601
+ .add_data(line["premise"])
602
+ .add_data(line["hypothesis"])
603
+ .add_end(
604
+ (
605
+ "Réponds uniquement par :\n"
606
+ "0 - si la deuxième phrase implique la première,\n"
607
+ "1 - si la relation est neutre,\n"
608
+ "2 - s'il y a contradiction.\n"
609
+ "Réponds uniquement par 0, 1 ou 2. La réponse est :"
610
+ )
611
+ )
612
+ .build(),
613
+ line_to_data_fn=lambda line: {
614
+ "premise": line["premise"],
615
+ "hypothesis": line["hypothesis"],
616
+ },
617
+ ),
618
+ BorealTasks.TIMELINE.value: Dataset(
619
+ name=BorealTasks.TIMELINE.value,
620
+ description="Binary temporal ordering task: predict which of two events happened first.",
621
+ possible_ground_truths=["1", "2"],
622
+ hugging_face_repo=COLE_REPOSITORY_NAME,
623
+ line_to_truth_fn=lambda line: str(line["answer"]),
624
+ line_to_prompt_fn=lambda line: (
625
+ PromptBuilder()
626
+ .add_premise(
627
+ "Deux événements historiques sont donnés. Lequel est survenu en premier ?"
628
+ )
629
+ .add_data(f"Événement 1 : {line['event_1']} (date: {line['date_1']})")
630
+ .add_data(f"Événement 2 : {line['event_2']} (date: {line['date_2']})")
631
+ .add_end(
632
+ "Réponds uniquement par 1 si l'événement 1 est venu avant l'événement 2, "
633
+ "ou 2 si l'événement 2 est venu avant l'événement 1. La réponse est :"
634
+ )
635
+ .build()
636
+ ),
637
+ line_to_data_fn=lambda line: {
638
+ "event_1": line["event_1"],
639
+ "date_1": line["date_1"],
640
+ "event_2": line["event_2"],
641
+ "date_2": line["date_2"],
642
+ },
643
+ ),
644
+ BorealTasks.LQLE.value: Dataset(
645
+ name=BorealTasks.LQLE.value,
646
+ description=(
647
+ "Author name prediction for Quebec literature works: "
648
+ "given the title of a literary work, predict the Quebecois author name."
649
+ ),
650
+ possible_ground_truths=[],
651
+ hugging_face_repo=COLE_REPOSITORY_NAME,
652
+ line_to_truth_fn=lambda line: line["author"],
653
+ line_to_prompt_fn=lambda line: (
654
+ PromptBuilder()
655
+ .add_premise(
656
+ "On te donne le titre d'une oeuvre de littérature québécoise. "
657
+ "Ton rôle est de donner UNIQUEMENT le nom de l'auteur ou de l'autrice québécois(e) qui l'a écrit."
658
+ )
659
+ .add_data(f"Titre : {line.get('work_title', '').strip()}")
660
+ .add_end(
661
+ "Réponds uniquement par le nom complet de l'auteur ou de l'autrice, sans commentaire ni guillemets. "
662
+ "La réponse est :"
663
+ )
664
+ .build()
665
+ ),
666
+ line_to_data_fn=lambda line: {
667
+ "work_title": line["work_title"],
668
+ },
669
+ ),
670
+ BorealTasks.PIQAQFR.value: Dataset(
671
+ name=BorealTasks.PIQAQFR.value,
672
+ description=(
673
+ "Physical commonsense multiple-choice task (Global PIQA) in Quebec French: "
674
+ "given a short situation and two possible solutions, choose the physically plausible one."
675
+ ),
676
+ possible_ground_truths=["0", "1"],
677
+ hugging_face_repo=COLE_REPOSITORY_NAME,
678
+ line_to_truth_fn=lambda line: str(line["label"]),
679
+ line_to_prompt_fn=lambda line: (
680
+ PromptBuilder()
681
+ .add_premise(
682
+ "Lis la situation suivante et choisis l'action qui a le plus de sens sur le plan physique."
683
+ )
684
+ .add_data(f"Problème : {line['prompt']}")
685
+ .add_data(f"Option 0 : {line['solution0']}")
686
+ .add_data(f"Option 1 : {line['solution1']}")
687
+ .add_end(
688
+ "Réponds uniquement par 0 si l'option 0 est la meilleure, "
689
+ "ou par 1 si l'option 1 est la meilleure. La réponse est :"
690
+ )
691
+ .build()
692
+ ),
693
+ line_to_data_fn=lambda line: {
694
+ "prompt": line["prompt"],
695
+ "solution0": line["solution0"],
696
+ "solution1": line["solution1"],
697
+ },
698
+ ),
699
+ BorealTasks.PIQAFR.value: Dataset(
700
+ name=BorealTasks.PIQAFR.value,
701
+ description=(
702
+ "Physical commonsense multiple-choice task (Global PIQA) in standard French: "
703
+ "given a short situation and two possible solutions, choose the physically plausible one."
704
+ ),
705
+ possible_ground_truths=["0", "1"],
706
+ hugging_face_repo=COLE_REPOSITORY_NAME,
707
+ line_to_truth_fn=lambda line: str(line["label"]),
708
+ line_to_prompt_fn=lambda line: (
709
+ PromptBuilder()
710
+ .add_premise(
711
+ "Lis la situation suivante et choisis l'action qui a le plus de sens sur le plan physique."
712
+ )
713
+ .add_data(f"Problème : {line['prompt']}")
714
+ .add_data(f"Option 0 : {line['solution0']}")
715
+ .add_data(f"Option 1 : {line['solution1']}")
716
+ .add_end(
717
+ "Réponds uniquement par 0 si l'option 0 est la meilleure, "
718
+ "ou par 1 si l'option 1 est la meilleure. La réponse est :"
719
+ )
720
+ .build()
721
+ ),
722
+ line_to_data_fn=lambda line: {
723
+ "prompt": line["prompt"],
724
+ "solution0": line["solution0"],
725
+ "solution1": line["solution1"],
726
+ },
727
+ ),
728
+ BorealTasks.QCCR.value: Dataset(
729
+ name=BorealTasks.QCCR.value,
730
+ description="Québec cities – predict the administrative region of a given city.",
731
+ possible_ground_truths=[],
732
+ hugging_face_repo=COLE_REPOSITORY_NAME,
733
+ line_to_truth_fn=lambda line: line["region"],
734
+ line_to_prompt_fn=lambda line: (
735
+ PromptBuilder()
736
+ .add_premise(
737
+ "À quelle région administrative appartient cette ville du Québec ?"
738
+ )
739
+ .add_data(f"Ville : {line['city']}")
740
+ .add_end(
741
+ "Réponds uniquement par le nom exact de la région. La réponse est :"
742
+ )
743
+ .build()
744
+ ),
745
+ line_to_data_fn=lambda line: {
746
+ "city": line["city"],
747
+ },
748
+ ),
749
+ BorealTasks.QCCY.value: Dataset(
750
+ name=BorealTasks.QCCY.value,
751
+ description="Québec cities – predict the year the city was founded.",
752
+ possible_ground_truths=[],
753
+ hugging_face_repo=COLE_REPOSITORY_NAME,
754
+ line_to_truth_fn=lambda line: line["founded_year"],
755
+ line_to_prompt_fn=lambda line: (
756
+ PromptBuilder()
757
+ .add_premise("En quelle année cette ville du Québec a-t-elle été fondée ?")
758
+ .add_data(f"Ville : {line['city']}")
759
+ .add_end("Réponds uniquement par l'année (4 chiffres). La réponse est :")
760
+ .build()
761
+ ),
762
+ line_to_data_fn=lambda line: {
763
+ "city": line["city"],
764
+ },
765
+ ),
766
+ BorealTasks.QCCP.value: Dataset(
767
+ name=BorealTasks.QCCP.value,
768
+ description="Québec cities – predict the population of the city (estimations de population, année 2024)",
769
+ possible_ground_truths=[],
770
+ hugging_face_repo=COLE_REPOSITORY_NAME,
771
+ line_to_truth_fn=lambda line: line["population"],
772
+ line_to_prompt_fn=lambda line: (
773
+ PromptBuilder()
774
+ .add_premise("Quelle est la population de cette ville du Québec en 2024 ?")
775
+ .add_data(f"Ville : {line['city']}")
776
+ .add_end("Réponds uniquement par le nombre d'habitants. La réponse est :")
777
+ .build()
778
+ ),
779
+ line_to_data_fn=lambda line: {
780
+ "city": line["city"],
781
+ },
782
+ ),
783
+ }
784
+
785
+
786
+ def preload_all_datasets():
787
+ """Loads all datasets into cache for later usage"""
788
+ for dataset in datasets.values():
789
+ dataset.load_data()
790
+
791
+
792
+ def generate_metadata_dict():
793
+ """Generates a dictionary with all the datasets metadata information"""
794
+ metadata_dict = {}
795
+ for dataset in datasets.values():
796
+ metadata_dict[dataset.name] = dataset.metadata
797
+ return metadata_dict
cole/dataset/prompt_builder.py ADDED
@@ -0,0 +1,43 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import logging
2
+ from typing import List
3
+
4
+
5
+ class PromptBuilder:
6
+ """Builder class for creating prompt strings with dynamic data."""
7
+
8
+ def __init__(self):
9
+ self.premise: List[str] = []
10
+ self.end: List[str] = []
11
+ self.data: List[str] = []
12
+ self.data_only = False
13
+
14
+ def add_data(self, data):
15
+ self.data.append(data)
16
+ return self
17
+
18
+ def add_end(self, end):
19
+ self.end.append(end)
20
+ return self
21
+
22
+ def set_data_only(self, data_only):
23
+ self.data_only = data_only
24
+ return self
25
+
26
+ def add_premise(self, premise):
27
+ self.premise.append(premise)
28
+ return self
29
+
30
+ def build(self):
31
+ """Builds and returns the prompt as a string based on data, premise and end that were added to the builder."""
32
+ if len(self.data) == 0:
33
+ logging.warning(
34
+ "This prompt did not contain any data, was that intentional ?"
35
+ )
36
+
37
+ data = "\n".join(self.data)
38
+ if self.data_only:
39
+ return data
40
+
41
+ end = "".join(self.end)
42
+ premise = "".join(self.premise)
43
+ return f"{premise}\n{data}\n{end}"
cole/docker_requirements.txt ADDED
@@ -0,0 +1,18 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Core
2
+ python-dotenv
3
+ numpy
4
+ python-multipart
5
+ scikit-learn
6
+ # Web backend (si tu utilises FastAPI/Flask, sinon à ignorer)
7
+ fastapi
8
+ uvicorn
9
+ slowapi
10
+ # Optional: pretty printing, progress bars, etc.
11
+ tqdm
12
+ aenum
13
+
14
+ evaluate
15
+ wheel
16
+
17
+ # Pour compatibilité ancienne
18
+ protobuf<=7.36.0
cole/evaluation/__init__.py ADDED
File without changes
cole/evaluation/evaluation_pipeline.py ADDED
@@ -0,0 +1,159 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import argparse
2
+ import gc
3
+ import logging
4
+ from datetime import datetime
5
+
6
+ import torch
7
+ import wandb
8
+ from tqdm import tqdm
9
+
10
+ from predictions.all_llms import llms
11
+ from cole.evaluation.llm_evaluator import ModelEvaluator
12
+ from cole.evaluation.llm_factory import model_factory
13
+ from cole.evaluation.tools import split_llm_list, str2bool
14
+ from cole.task.task_factory import tasks_factory
15
+ from cole.task.task_names import COLETasks, BorealTasks
16
+
17
+ parser = argparse.ArgumentParser()
18
+ parser.add_argument(
19
+ "--test",
20
+ help="If set to true, the system will default to testing only a small model with a few examples.",
21
+ default=False,
22
+ type=str2bool,
23
+ )
24
+ parser.add_argument(
25
+ "--max_examples",
26
+ "-m",
27
+ help="The maximum number of examples to use, defaults to None.",
28
+ type=int,
29
+ default=None,
30
+ )
31
+ parser.add_argument(
32
+ "--models_name",
33
+ "-mn",
34
+ help="The name of the model(s) to load.",
35
+ type=str,
36
+ default=None,
37
+ )
38
+
39
+ parser.add_argument(
40
+ "--batch_size",
41
+ help="The batch size to use during the evaluation.",
42
+ type=int,
43
+ default=32,
44
+ )
45
+
46
+ parser.add_argument(
47
+ "--llm_split",
48
+ help="The split of the LLMs list to use. It can be '1', '2' or '3'.",
49
+ type=int,
50
+ default=None,
51
+ choices=[1, 2, 3],
52
+ )
53
+
54
+ parser.add_argument(
55
+ "--skip_first_n",
56
+ help="The number of LLM to skip in the list of split",
57
+ type=int,
58
+ default=None,
59
+ )
60
+ parser.add_argument(
61
+ "--tasks_group",
62
+ help="The task group to test",
63
+ type=str,
64
+ default=None,
65
+ choices=["all", "cole", "boreal", "comparison"],
66
+ )
67
+
68
+ args = parser.parse_args()
69
+
70
+ if args.tasks_group == "all":
71
+ tasks_names = list(COLETasks) + list(BorealTasks)
72
+ from cole import complete as project
73
+ elif args.tasks_group == "cole":
74
+ tasks_names = list(COLETasks)
75
+ from cole import cole as project
76
+ elif args.tasks_group == "boreal":
77
+ tasks_names = list(BorealTasks)
78
+ from cole import boreal as project
79
+ elif args.tasks_group == "comparison":
80
+ tasks_names = ["frcoe"]
81
+ from cole import comparison as project
82
+ else:
83
+ raise ValueError("Invalid value for tasks_group")
84
+
85
+ tasks = tasks_factory(tasks_names)
86
+
87
+ models = []
88
+ if args.models_name is not None:
89
+ if args.models_name in llms:
90
+ models = llms[args.models_name]
91
+ else:
92
+ models = args.models_name.split(",")
93
+ else:
94
+ models = llms["all"]
95
+
96
+ models = split_llm_list(models=models, llm_split=args.llm_split)
97
+
98
+ if args.skip_first_n is not None:
99
+ models = models[args.skip_first_n :]
100
+
101
+ logging.info("Starting Evaluation")
102
+
103
+ time_start = datetime.now()
104
+
105
+ for model_name in tqdm(
106
+ models, total=len(models), desc="Processing LLM inference on tasks."
107
+ ):
108
+ model = None
109
+ evaluator = None
110
+ try:
111
+ model = model_factory(model_name, batch_size=args.batch_size)
112
+ logging.info("Creating model")
113
+ evaluator = ModelEvaluator()
114
+ logging.info("Evaluating model")
115
+
116
+ exp_name = f"{model_name}"
117
+ wandb.init(
118
+ project=project,
119
+ entity="doctorate",
120
+ config={
121
+ "model_name": model_name,
122
+ "tasks": "; ".join(
123
+ t.value if hasattr(t, "value") else str(t) for t in tasks_names
124
+ ),
125
+ "batch_size": args.batch_size,
126
+ },
127
+ name=exp_name,
128
+ )
129
+
130
+ predictions_payload = evaluator.evaluate_subset(model, tasks, args.max_examples)
131
+ wandb.log(predictions_payload)
132
+
133
+ logging.info("Saving results")
134
+ evaluator.save_results("./results")
135
+
136
+ metrics_payload = evaluator.compute_metrics()
137
+ evaluator.save_metrics("./results")
138
+ wandb.log(metrics_payload)
139
+
140
+ except Exception as e:
141
+ error_message = f"Evaluation failed for model {model_name}: {e}"
142
+ logging.error(error_message)
143
+ wandb.finish(exit_code=1)
144
+ continue
145
+ else:
146
+ wandb.finish(exit_code=0)
147
+ finally:
148
+ # Memory cleaning
149
+ del model
150
+ del evaluator
151
+ gc.collect()
152
+ torch.cuda.empty_cache()
153
+
154
+ time_end = datetime.now()
155
+ info_message = f"End time: {time_end}"
156
+ logging.info(info_message)
157
+ elapsed_time = time_end - time_start
158
+ info_message = f"Elapsed time: {elapsed_time}"
159
+ logging.info(info_message)
cole/evaluation/evaluation_pipeline_private_llm.py ADDED
@@ -0,0 +1,137 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import argparse
2
+ import gc
3
+ import logging
4
+ from datetime import datetime
5
+
6
+ import wandb
7
+ from dotenv import load_dotenv
8
+ from tqdm import tqdm
9
+
10
+ from predictions.all_llms import private_llm
11
+ from cole.evaluation.llm_evaluator import ModelEvaluator
12
+ from cole.evaluation.tools import str2bool
13
+ from cole.language_model.private_lm import RemoteLLMModel
14
+ from cole.task.task_factory import tasks_factory
15
+ from cole.task.task_names import COLETasks, BorealTasks
16
+
17
+ load_dotenv(".env")
18
+
19
+ parser = argparse.ArgumentParser()
20
+ parser.add_argument(
21
+ "--test",
22
+ help="If set to true, the system will default to testing only a small model with a few examples.",
23
+ default=False,
24
+ type=str2bool,
25
+ )
26
+ parser.add_argument(
27
+ "--max_examples",
28
+ "-m",
29
+ help="The maximum number of examples to use, defaults to None.",
30
+ type=int,
31
+ default=None,
32
+ )
33
+ parser.add_argument(
34
+ "--models_name",
35
+ "-mn",
36
+ help="The name of the model(s) to load.",
37
+ type=str,
38
+ default=None,
39
+ )
40
+ parser.add_argument(
41
+ "--provider_name",
42
+ "-pn",
43
+ help="The name of the LLM provider to load.",
44
+ type=str,
45
+ default=None,
46
+ choices=list(private_llm.keys()),
47
+ )
48
+ parser.add_argument(
49
+ "--tasks_group",
50
+ help="The task group to test",
51
+ type=str,
52
+ default=None,
53
+ choices=["all", "cole", "boreal", "comparison"],
54
+ )
55
+
56
+ args = parser.parse_args()
57
+
58
+ if args.tasks_group == "all":
59
+ tasks_names = list(COLETasks) + list(BorealTasks)
60
+ from cole import complete as project
61
+ elif args.tasks_group == "cole":
62
+ tasks_names = list(COLETasks)
63
+ from cole import cole as project
64
+ elif args.tasks_group == "boreal":
65
+ tasks_names = list(BorealTasks)
66
+ from cole import boreal as project
67
+ elif args.tasks_group == "comparison":
68
+ tasks_names = ["frcoe"]
69
+ from cole import comparison as project
70
+ else:
71
+ raise ValueError("Invalid value for tasks_group")
72
+
73
+ tasks = tasks_factory(tasks_names)
74
+
75
+ models = []
76
+ if args.models_name is not None:
77
+ if args.models_name in private_llm:
78
+ models = private_llm[args.models_name]
79
+ else:
80
+ models = args.models_name.split(",")
81
+ elif args.provider_name is not None:
82
+ models = private_llm[args.provider_name]
83
+ else:
84
+ models = private_llm["all"]
85
+
86
+ logging.info("Starting Evaluation")
87
+
88
+ time_start = datetime.now()
89
+
90
+ for model_name in tqdm(
91
+ models, total=len(models), desc="Processing LLM inference on tasks."
92
+ ):
93
+ model = None
94
+ evaluator = None
95
+ try:
96
+ model = RemoteLLMModel(model_name=model_name)
97
+ logging.info("Creating model")
98
+ evaluator = ModelEvaluator()
99
+ logging.info("Evaluating model")
100
+
101
+ exp_name = f"{model_name}"
102
+ wandb.init(
103
+ project=project,
104
+ entity="doctorate",
105
+ config={"model_name": model_name, "tasks": "; ".join(tasks_names)},
106
+ name=exp_name,
107
+ )
108
+
109
+ predictions_payload = evaluator.evaluate_subset(model, tasks, args.max_examples)
110
+ logging.info("Writing predictions to WandB.")
111
+ wandb.log(predictions_payload)
112
+
113
+ logging.info("Saving results")
114
+ evaluator.save_results("./results")
115
+
116
+ metrics_payload = evaluator.compute_metrics()
117
+ evaluator.save_metrics("./results")
118
+ wandb.log(metrics_payload)
119
+
120
+ except Exception as e:
121
+ error_message = f"Evaluation failed for model {model_name}: {e}"
122
+ logging.error(error_message)
123
+ wandb.finish(exit_code=1)
124
+ continue
125
+ finally:
126
+ # Memory cleaning
127
+ del model
128
+ del evaluator
129
+ gc.collect()
130
+ wandb.finish(exit_code=0)
131
+
132
+ time_end = datetime.now()
133
+ info_message = f"End time: {time_end}"
134
+ logging.info(info_message)
135
+ elapsed_time = time_end - time_start
136
+ info_message = f"Elapsed time: {elapsed_time}"
137
+ logging.info(info_message)
cole/evaluation/evaluation_pipeline_small.py ADDED
@@ -0,0 +1,150 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import argparse
2
+ import gc
3
+ import logging
4
+ from datetime import datetime
5
+
6
+ import torch
7
+ import wandb
8
+ from tqdm import tqdm
9
+
10
+ from predictions.all_llms import small_llm
11
+ from cole.evaluation.llm_evaluator import ModelEvaluator
12
+ from cole.evaluation.llm_factory import model_factory
13
+ from cole.evaluation.tools import str2bool
14
+ from cole.task.task_factory import tasks_factory
15
+ from cole.task.task_names import COLETasks, BorealTasks
16
+
17
+ parser = argparse.ArgumentParser()
18
+ parser.add_argument(
19
+ "--test",
20
+ help="If set to true, the system will default to testing only a small model with a few examples.",
21
+ default=False,
22
+ type=str2bool,
23
+ )
24
+ parser.add_argument(
25
+ "--max_examples",
26
+ "-m",
27
+ help="The maximum number of examples to use, defaults to None.",
28
+ type=int,
29
+ default=None,
30
+ )
31
+ parser.add_argument(
32
+ "--token",
33
+ "-t",
34
+ help="Input your HuggingFace token to fetch models.",
35
+ type=str,
36
+ default=None,
37
+ )
38
+ parser.add_argument(
39
+ "--models_name",
40
+ "-mn",
41
+ help="The name of the model(s) to load.",
42
+ type=str,
43
+ default=None,
44
+ )
45
+
46
+ parser.add_argument(
47
+ "--batch_size",
48
+ help="The batch size to use during the evaluation.",
49
+ type=int,
50
+ default=32,
51
+ )
52
+
53
+ parser.add_argument(
54
+ "--skip_first_n",
55
+ help="The number of LLM to skip in the list of split",
56
+ type=int,
57
+ default=None,
58
+ )
59
+
60
+ parser.add_argument(
61
+ "--tasks_group",
62
+ help="The task group to test",
63
+ type=str,
64
+ default=None,
65
+ choices=["all", "cole", "boreal", "comparison"],
66
+ )
67
+
68
+ args = parser.parse_args()
69
+
70
+ if args.tasks_group == "all":
71
+ tasks_names = list(COLETasks) + list(BorealTasks)
72
+ from cole import complete as project
73
+ elif args.tasks_group == "cole":
74
+ tasks_names = list(COLETasks)
75
+ from cole import cole as project
76
+ elif args.tasks_group == "boreal":
77
+ tasks_names = list(BorealTasks)
78
+ from cole import boreal as project
79
+ elif args.tasks_group == "comparison":
80
+ tasks_names = ["frcoe"]
81
+ from cole import comparison as project
82
+ else:
83
+ raise ValueError("Invalid value for tasks_group")
84
+
85
+ tasks = tasks_factory(tasks_names)
86
+
87
+ models = []
88
+ if args.models_name is not None:
89
+ if args.models_name in small_llm:
90
+ models = small_llm[args.models_name]
91
+ else:
92
+ models = args.models_name.split(",")
93
+ else:
94
+ models = small_llm["all"]
95
+
96
+ if args.skip_first_n is not None:
97
+ models = models[args.skip_first_n :]
98
+
99
+ logging.info("Starting Evaluation")
100
+
101
+ time_start = datetime.now()
102
+
103
+ for model_name in tqdm(
104
+ models, total=len(models), desc="Processing LLM inference on tasks."
105
+ ):
106
+ model = None
107
+ evaluator = None
108
+ try:
109
+ model = model_factory(model_name, batch_size=args.batch_size)
110
+ logging.info("Creating model")
111
+ evaluator = ModelEvaluator()
112
+ logging.info("Evaluating model")
113
+
114
+ exp_name = f"{model_name}"
115
+ wandb.init(
116
+ project=project,
117
+ entity="doctorate",
118
+ config={"model_name": model_name, "tasks": "; ".join(tasks_names)},
119
+ name=exp_name,
120
+ )
121
+
122
+ predictions_payload = evaluator.evaluate_subset(model, tasks, args.max_examples)
123
+ wandb.log(predictions_payload)
124
+
125
+ logging.info("Saving results")
126
+ evaluator.save_results("./results")
127
+
128
+ metrics_payload = evaluator.compute_metrics()
129
+ evaluator.save_metrics("./results")
130
+ wandb.log(metrics_payload)
131
+
132
+ except Exception as e:
133
+ error_message = f"Evaluation failed for model {model_name}: {e}"
134
+ logging.error(error_message)
135
+ wandb.finish(exit_code=1)
136
+ continue
137
+ finally:
138
+ # Memory cleaning
139
+ del model
140
+ del evaluator
141
+ gc.collect()
142
+ torch.cuda.empty_cache()
143
+ wandb.finish(exit_code=0)
144
+
145
+ time_end = datetime.now()
146
+ info_message = f"End time: {time_end}"
147
+ logging.info(info_message)
148
+ elapsed_time = time_end - time_start
149
+ info_message = f"Elapsed time: {elapsed_time}"
150
+ logging.info(info_message)
cole/evaluation/evaluation_pipeline_small_2.py ADDED
@@ -0,0 +1,151 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import argparse
2
+ import gc
3
+ import logging
4
+ from datetime import datetime
5
+
6
+ import torch
7
+ import wandb
8
+ from tqdm import tqdm
9
+
10
+ from predictions.all_llms import small_llm_2
11
+ from cole.evaluation.llm_evaluator import ModelEvaluator
12
+ from cole.evaluation.llm_factory import model_factory
13
+ from cole.evaluation.tools import str2bool
14
+ from cole.task.task_factory import tasks_factory
15
+ from cole.task.task_names import COLETasks, BorealTasks
16
+
17
+ parser = argparse.ArgumentParser()
18
+ parser.add_argument(
19
+ "--test",
20
+ help="If set to true, the system will default to testing only a small model with a few examples.",
21
+ default=False,
22
+ type=str2bool,
23
+ )
24
+ parser.add_argument(
25
+ "--max_examples",
26
+ "-m",
27
+ help="The maximum number of examples to use, defaults to None.",
28
+ type=int,
29
+ default=None,
30
+ )
31
+ parser.add_argument(
32
+ "--token",
33
+ "-t",
34
+ help="Input your HuggingFace token to fetch models.",
35
+ type=str,
36
+ default=None,
37
+ )
38
+ parser.add_argument(
39
+ "--models_name",
40
+ "-mn",
41
+ help="The name of the model(s) to load.",
42
+ type=str,
43
+ default=None,
44
+ )
45
+
46
+ parser.add_argument(
47
+ "--batch_size",
48
+ help="The batch size to use during the evaluation.",
49
+ type=int,
50
+ default=256,
51
+ )
52
+
53
+ parser.add_argument(
54
+ "--skip_first_n",
55
+ help="The number of LLM to skip in the list of split",
56
+ type=int,
57
+ default=None,
58
+ )
59
+ parser.add_argument(
60
+ "--tasks_group",
61
+ help="The task group to test",
62
+ type=str,
63
+ default=None,
64
+ choices=["all", "cole", "boreal", "comparison"],
65
+ )
66
+
67
+ args = parser.parse_args()
68
+
69
+ if args.tasks_group == "all":
70
+ tasks_names = list(COLETasks) + list(BorealTasks)
71
+ from cole import complete as project
72
+ elif args.tasks_group == "cole":
73
+ tasks_names = list(COLETasks)
74
+ from cole import cole as project
75
+ elif args.tasks_group == "boreal":
76
+ tasks_names = list(BorealTasks)
77
+ from cole import boreal as project
78
+ elif args.tasks_group == "comparison":
79
+ tasks_names = ["frcoe"]
80
+ from cole import comparison as project
81
+ else:
82
+ raise ValueError("Invalid value for tasks_group")
83
+
84
+ tasks = tasks_factory(tasks_names)
85
+
86
+ models = []
87
+ if args.models_name is not None:
88
+ if args.models_name in small_llm_2:
89
+ models = small_llm_2[args.models_name]
90
+ elif args.models_name == "RandomBaselineModel":
91
+ models = ["RandomBaselineModel"]
92
+ else:
93
+ models = args.models_name.split(",")
94
+ else:
95
+ models = small_llm_2["all"]
96
+
97
+ if args.skip_first_n is not None:
98
+ models = models[args.skip_first_n :]
99
+
100
+ logging.info("Starting Evaluation")
101
+
102
+ time_start = datetime.now()
103
+
104
+ for model_name in tqdm(
105
+ models, total=len(models), desc="Processing LLM inference on tasks."
106
+ ):
107
+ model = None
108
+ evaluator = None
109
+ try:
110
+ model = model_factory(model_name, batch_size=args.batch_size)
111
+ logging.info("Creating model")
112
+ evaluator = ModelEvaluator()
113
+ logging.info("Evaluating model")
114
+
115
+ exp_name = f"{model_name}"
116
+ wandb.init(
117
+ project=project,
118
+ entity="doctorate",
119
+ config={"model_name": model_name, "tasks": "; ".join(tasks_names)},
120
+ name=exp_name,
121
+ )
122
+
123
+ predictions_payload = evaluator.evaluate_subset(model, tasks, args.max_examples)
124
+ wandb.log(predictions_payload)
125
+
126
+ logging.info("Saving results")
127
+ evaluator.save_results("./results")
128
+
129
+ metrics_payload = evaluator.compute_metrics()
130
+ evaluator.save_metrics("./results")
131
+ wandb.log(metrics_payload)
132
+
133
+ except Exception as e:
134
+ error_message = f"Evaluation failed for model {model_name}: {e}"
135
+ logging.error(error_message)
136
+ wandb.finish(exit_code=1)
137
+ continue
138
+ finally:
139
+ # Memory cleaning
140
+ del model
141
+ del evaluator
142
+ gc.collect()
143
+ torch.cuda.empty_cache()
144
+ wandb.finish(exit_code=0)
145
+
146
+ time_end = datetime.now()
147
+ info_message = f"End time: {time_end}"
148
+ logging.info(info_message)
149
+ elapsed_time = time_end - time_start
150
+ info_message = f"Elapsed time: {elapsed_time}"
151
+ logging.info(info_message)
cole/evaluation/llm_evaluator.py ADDED
@@ -0,0 +1,217 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import gc
2
+ import json
3
+ import logging
4
+ import os
5
+ from datetime import datetime
6
+ from typing import Dict, List
7
+
8
+ import torch
9
+ import wandb
10
+ from datasets import Dataset
11
+ from tqdm import tqdm
12
+
13
+ from cole.language_model.language_model_abstraction import LanguageModel
14
+ from cole.task.task import Task
15
+ from cole.task.task_factory import tasks_factory
16
+
17
+
18
+ class ModelEvaluator:
19
+ """
20
+ The model evaluator acts as a pipeline for evaluating models on tasks available from tasks_factory.
21
+ """
22
+
23
+ def __init__(self):
24
+ self.last_predictions = {}
25
+
26
+ self.last_metrics = {}
27
+
28
+ self.last_model_name = None
29
+
30
+ def compute_metrics(self) -> Dict:
31
+ """
32
+ Compute metrics over the last tested model's predictions,
33
+ must have called one the evaluate functions before or loaded predictions with load_predictions_from_file.
34
+ """
35
+ task_metrics: List[Dict] = []
36
+
37
+ for task_dict in tqdm(
38
+ self.last_predictions["tasks"], desc="Computing metrics: "
39
+ ):
40
+ task_name, preds = list(task_dict.items())[0]
41
+ if not preds:
42
+ warning_message = f"Task '{task_name}' ignored due to no predictions"
43
+ logging.warning(warning_message)
44
+ continue
45
+ try:
46
+ tasks = tasks_factory([task_name])
47
+ task = tasks[0]
48
+ metric_score, warning = task.compute(preds)
49
+ except Exception as e:
50
+ error_message = f"Error while calculating metrics'{task_name}' : {e}"
51
+ logging.error(error_message)
52
+ continue
53
+ metric_name = task.metric_name
54
+ task_entry = {
55
+ task_name: {
56
+ metric_name: {**metric_score, f"{metric_name}_warning": warning}
57
+ }
58
+ }
59
+ task_metrics.append(task_entry)
60
+ wandb.log({f"{task_name}.{metric_name}": {**metric_score}})
61
+
62
+ # Build the response from a fresh dict so `last_metrics` does NOT alias
63
+ # `last_predictions`. The previous implementation stored the same dict
64
+ # reference in both, so mutating one would silently mutate the other.
65
+ self.last_metrics = {
66
+ **{k: v for k, v in self.last_predictions.items() if k != "tasks"},
67
+ "tasks": task_metrics,
68
+ }
69
+ return self.last_metrics
70
+
71
+ def load_predictions_from_file(self, file_path: str) -> None:
72
+ """
73
+ Load predictions from file to compute metrics :param file_path:path to the predictions file.
74
+ """
75
+
76
+ try:
77
+ with open(file_path, "r", encoding="utf-8") as f:
78
+ self.last_predictions = json.load(f)
79
+ except FileNotFoundError:
80
+ error = f"File not found: {file_path}"
81
+ logging.error(error)
82
+ self.last_predictions = None
83
+ except json.JSONDecodeError:
84
+ error = f"Invalid JSON in file: {file_path}"
85
+ logging.error(error)
86
+ self.last_predictions = None
87
+
88
+ def save_metrics(self, save_path):
89
+ """
90
+ Saves computed metrics to a json file.
91
+ :param save_path : the path to which the json file will be saved.
92
+ """
93
+ # `last_metrics` is initialized to {} (not None), so the previous
94
+ # `is None` check never fired. Treat both empty/missing the same way.
95
+ if not self.last_metrics or self.last_model_name is None:
96
+ logging.info("No metrics saved")
97
+ return None
98
+ return self.save_object(
99
+ save_path,
100
+ self.last_metrics,
101
+ f"{self.last_model_name.replace('/', '_')}_metrics.json",
102
+ )
103
+
104
+ def evaluate(self, model: LanguageModel, tasks: List[Task]):
105
+ """
106
+ Evaluates a given model on the given tasks.
107
+ :param model : the model that will infer on the given tasks.
108
+ :param tasks : the tasks to be evaluated on.
109
+ """
110
+ return self.evaluate_subset(model, tasks)
111
+
112
+ def evaluate_subset(
113
+ self, model: LanguageModel, tasks: List[Task], subset_size=None
114
+ ) -> Dict:
115
+ """
116
+ Evaluates a given model on the given tasks, but only on a given size.
117
+ :param model : the model that will infer on the given tasks.
118
+ :param tasks : the tasks to be evaluated on.
119
+ :param subset_size : the size of the subset to be evaluated.
120
+ """
121
+ predictions = []
122
+ for task in tasks:
123
+ info_log = (
124
+ f"-----Doing task '{task.task_name}' with model '{model.name}-----'."
125
+ )
126
+ logging.info(info_log)
127
+ # Initialize before try so the finally block can drop them without
128
+ # relying on `locals()` introspection (which is brittle and was
129
+ # papering over the case where the try block raised before assignment).
130
+ prompts = None
131
+ evaluate_dataset = None
132
+ try:
133
+ if subset_size is None:
134
+ prompts = task.dataset.prompts[:]
135
+ else:
136
+ prompts = task.dataset.prompts[:subset_size]
137
+
138
+ evaluate_dataset = Dataset.from_dict({"text": prompts})
139
+
140
+ task_predictions = model.predict(
141
+ evaluation_dataset=evaluate_dataset, task=task
142
+ )
143
+
144
+ task_predictions = {task.task_name: task_predictions}
145
+ predictions.append(task_predictions)
146
+
147
+ except Exception as e:
148
+ error_message = f"Task '{task.task_name}' has failed : {e}"
149
+ logging.error(error_message)
150
+ wandb.log({task.task_name: "Failed"})
151
+ continue
152
+ finally:
153
+ # Memory cleaning
154
+ del evaluate_dataset
155
+ del prompts
156
+ torch.cuda.empty_cache()
157
+ gc.collect()
158
+ wandb.log({task.task_name: "Success"})
159
+ logging.info("Finished evaluating tasks.")
160
+ self.last_predictions = {
161
+ "model_name": model.name,
162
+ "model_url": f"https://huggingface.co/{model.name}",
163
+ "tasks": predictions,
164
+ }
165
+ self.last_model_name = model.name
166
+ return self.last_predictions
167
+
168
+ def save_results(self, save_path):
169
+ """
170
+ Saves inferred metrics to a json file.
171
+ :param save_path : the path to which the json file will be saved.
172
+ """
173
+
174
+ if self.last_model_name is None:
175
+ logging.error("Please evaluate before saving results")
176
+ return None
177
+ date_time_stamp = datetime.now().strftime("%Y%m%d-%H%M")
178
+ return self.save_object(
179
+ save_path,
180
+ self.last_predictions,
181
+ f"{self.last_model_name.replace('/', '_')}_{date_time_stamp}.json",
182
+ )
183
+
184
+ def save_object(self, save_dir_path, saved_object, filename):
185
+ """
186
+ Utility method to save the given object into a json file.
187
+ """
188
+ os.makedirs(save_dir_path, exist_ok=True)
189
+ full_path = os.path.join(save_dir_path, filename)
190
+ if os.path.isfile(full_path):
191
+ logging.info("Appending results to previous results file.")
192
+ try:
193
+ with open(full_path, "r", encoding="utf-8") as f:
194
+ data = json.load(f)
195
+ # `data.get("tasks")` could be None on a malformed/legacy file,
196
+ # in which case `.extend(...)` would raise AttributeError.
197
+ # Likewise the new payload may have no "tasks" key.
198
+ existing_tasks = data.get("tasks") or []
199
+ new_tasks = saved_object.get("tasks") or []
200
+ data["tasks"] = existing_tasks + new_tasks
201
+ with open(full_path, "w", encoding="utf-8") as f:
202
+ json.dump(data, f, indent=2)
203
+ info_message = f"Results saved to {save_dir_path}"
204
+ logging.info(info_message)
205
+ except Exception as e:
206
+ error_message = f"Failed to save object: {e}"
207
+ logging.error(error_message)
208
+ else:
209
+ try:
210
+ with open(full_path, "w", encoding="utf-8") as f:
211
+ json.dump(saved_object, f, indent=2)
212
+ info_message = f"Results saved to {save_dir_path}"
213
+ logging.info(info_message)
214
+ except Exception as e:
215
+ error_message = f"Failed to save object: {e}"
216
+ logging.error(error_message)
217
+ return full_path
cole/evaluation/llm_factory.py ADDED
@@ -0,0 +1,25 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from typing import Union
2
+
3
+ from predictions.all_llms import small_llm, llms, private_llm, small_llm_2
4
+ from cole.language_model.baseline import RandomBaselineModel
5
+ from cole.language_model.hugging_face_lm import HFLLMModel
6
+ from cole.language_model.language_model_abstraction import LanguageModel
7
+ from cole.language_model.private_lm import RemoteLLMModel
8
+
9
+
10
+ def model_factory(
11
+ model_name: str, batch_size: Union[int, None] = None
12
+ ) -> LanguageModel:
13
+ if model_name == "RandomBaselineModel":
14
+ model = RandomBaselineModel(model_name="random_baseline")
15
+ elif model_name in private_llm["all"]:
16
+ model = RemoteLLMModel(model_name=model_name)
17
+ elif (
18
+ model_name in llms["all"]
19
+ or model_name in small_llm["all"]
20
+ or model_name in small_llm_2["all"]
21
+ ):
22
+ model = HFLLMModel(model_name=model_name, batch_size=batch_size)
23
+ else:
24
+ raise ValueError(f"Model {model_name} not supported.")
25
+ return model
cole/evaluation/tools.py ADDED
@@ -0,0 +1,33 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import argparse
2
+ from typing import List, Union
3
+
4
+
5
+ def split_llm_list(models: List, llm_split: Union[None, int]) -> List:
6
+ if llm_split is None:
7
+ return models
8
+ if llm_split not in (1, 2, 3):
9
+ raise ValueError("llm_split must be in [1, 2, 3].")
10
+ if llm_split == 1:
11
+ models = models[: len(models) // 3]
12
+ elif llm_split == 2:
13
+ models = models[len(models) // 3 : 2 * len(models) // 3]
14
+ elif llm_split == 3:
15
+ models = models[2 * len(models) // 3 :]
16
+ return models
17
+
18
+
19
+ def str2bool(value: Union[bool, str]) -> bool:
20
+ """argparse-friendly bool parser.
21
+
22
+ `argparse(type=bool)` is a footgun: `bool("False")` is True, so any
23
+ non-empty string flips the flag to True. This converter rejects garbage
24
+ explicitly and accepts the obvious truthy/falsy spellings.
25
+ """
26
+ if isinstance(value, bool):
27
+ return value
28
+ normalized = str(value).strip().lower()
29
+ if normalized in ("true", "1", "yes", "y", "t"):
30
+ return True
31
+ if normalized in ("false", "0", "no", "n", "f"):
32
+ return False
33
+ raise argparse.ArgumentTypeError(f"Boolean value expected, got: {value!r}")
cole/language_model/__init__.py ADDED
File without changes
cole/language_model/anthropic_wrapper.py ADDED
@@ -0,0 +1,20 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from typing import Dict
2
+
3
+ from anthropic import Anthropic
4
+
5
+ from cole.language_model.open_ai_api_lm_wrapper import OpenAIAPILMWrapper
6
+
7
+
8
+ class AnthropicWrapper(OpenAIAPILMWrapper):
9
+ def __init__(self, model_name: str, api_key: str, extra_params: Dict):
10
+ super().__init__(model_name=model_name, extra_params=extra_params)
11
+ self.client = Anthropic(api_key=api_key)
12
+
13
+ self.tool_choices = {"type": "tool", "name": "classification"}
14
+
15
+ def _inner_generate_fn(self, prompt: Dict):
16
+ return self.client.messages.create(
17
+ model=self.model_name,
18
+ messages=prompt,
19
+ **self._extra_params,
20
+ )
cole/language_model/baseline.py ADDED
@@ -0,0 +1,45 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import numpy as np
2
+ from datasets.formatting.formatting import LazyRow
3
+
4
+ from cole.language_model.language_model_abstraction import LanguageModel
5
+ from cole.task.task import Task
6
+
7
+
8
+ class RandomBaselineModel(LanguageModel):
9
+ def __init__(self, model_name: str, seed: int = 42):
10
+ super().__init__(model_name)
11
+ self.random_generator = np.random.RandomState(seed=seed)
12
+
13
+ def predict(self, evaluation_dataset, task: Task):
14
+ size = len(evaluation_dataset)
15
+ choices = task.dataset.possible_ground_truths
16
+ if len(choices) == 0:
17
+ # Meaning it is a generation task.
18
+ predictions = []
19
+ for row in evaluation_dataset:
20
+
21
+ # Since we work with the prompt, including instruction, we extract the instance sentence after the
22
+ # "Phrase :", then we remove the trailing Part-of-speech content.
23
+ instance_sentence = row["text"].split("Phrase : ")[-1].split("\n")[0]
24
+ whitespace_instance_sentence = instance_sentence.split(" ")
25
+
26
+ choices = range(0, len(whitespace_instance_sentence))
27
+ prediction_idx = self.random_generator.choice(choices, size=1).tolist()[
28
+ 0
29
+ ]
30
+ prediction = whitespace_instance_sentence[prediction_idx]
31
+ predictions.append(prediction)
32
+ else:
33
+ # Meaning it is an inference task.
34
+ predictions = self.random_generator.choice(choices, size=size).tolist()
35
+ return predictions
36
+
37
+ def infer(self, rows: LazyRow) -> LazyRow:
38
+ return rows
39
+
40
+ def generate(self, rows: LazyRow) -> LazyRow:
41
+ return rows
42
+
43
+ @property
44
+ def num_parameters(self):
45
+ return 0
cole/language_model/cohere_wrapper.py ADDED
@@ -0,0 +1,31 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from typing import Dict, List, Union
2
+
3
+ from openai import OpenAI
4
+
5
+ from cole.language_model.open_ai_api_lm_wrapper import OpenAIAPILMWrapper
6
+
7
+
8
+ class CohereWrapper(OpenAIAPILMWrapper):
9
+ def __init__(
10
+ self,
11
+ model_name: str,
12
+ api_key: str,
13
+ extra_params: Dict,
14
+ use_function_calling: bool,
15
+ base_url: Union[None, str] = None,
16
+ timeout: int = 480,
17
+ ):
18
+ super().__init__(
19
+ model_name=model_name,
20
+ extra_params=extra_params,
21
+ use_function_calling=use_function_calling,
22
+ )
23
+ self.client = OpenAI(api_key=api_key, base_url=base_url, timeout=timeout)
24
+
25
+ def _inner_generate_fn(self, prompt: List):
26
+ return self.client.chat.completions.create(
27
+ model=self.model_name,
28
+ messages=prompt,
29
+ stream=False,
30
+ **self._extra_params,
31
+ )
cole/language_model/deepseek_wrapper.py ADDED
@@ -0,0 +1,30 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from typing import Dict, List
2
+
3
+ from openai import OpenAI
4
+
5
+ from cole.language_model.open_ai_api_lm_wrapper import OpenAIAPILMWrapper
6
+
7
+
8
+ class DeepSeekWrapper(OpenAIAPILMWrapper):
9
+ def __init__(
10
+ self,
11
+ model_name: str,
12
+ api_key: str,
13
+ extra_params: Dict,
14
+ use_function_calling: bool,
15
+ ):
16
+ super().__init__(
17
+ model_name=model_name,
18
+ extra_params=extra_params,
19
+ use_function_calling=use_function_calling,
20
+ )
21
+ self.client = OpenAI(api_key=api_key, base_url="https://api.deepseek.com")
22
+ self.tool_choices = "auto"
23
+
24
+ def _inner_generate_fn(self, prompt: List):
25
+ return self.client.chat.completions.create(
26
+ model=self.model_name,
27
+ messages=prompt,
28
+ stream=False,
29
+ **self._extra_params,
30
+ )
cole/language_model/google_wrapper.py ADDED
@@ -0,0 +1,22 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from typing import Dict, List
2
+
3
+ from openai import OpenAI
4
+
5
+ from cole.language_model.open_ai_api_lm_wrapper import OpenAIAPILMWrapper
6
+
7
+
8
+ class GoogleWrapper(OpenAIAPILMWrapper):
9
+ def __init__(self, model_name: str, api_key: str, extra_params: Dict):
10
+ super().__init__(model_name=model_name, extra_params=extra_params)
11
+ self.client = OpenAI(
12
+ api_key=api_key,
13
+ base_url="https://generativelanguage.googleapis.com/v1beta/openai/",
14
+ )
15
+
16
+ def _inner_generate_fn(self, prompt: List):
17
+ return self.client.chat.completions.create(
18
+ model=self.model_name,
19
+ messages=prompt,
20
+ stream=False,
21
+ **self._extra_params,
22
+ )
cole/language_model/hugging_face_lm.py ADDED
@@ -0,0 +1,161 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from typing import Union, List, Dict
2
+
3
+ import torch
4
+ from datasets import Dataset
5
+ from datasets.formatting.formatting import LazyRow
6
+ from transformers import (
7
+ pipeline,
8
+ )
9
+
10
+ from cole.language_model.language_model_abstraction import LanguageModel
11
+ from cole.language_model.huggingface_language_model_factory import (
12
+ hugging_face_language_model_tokenizer_factory,
13
+ )
14
+ from cole.task.task import TaskType, Task
15
+
16
+
17
+ class HFLLMModel(LanguageModel):
18
+ """
19
+ LLM Model based on Hugging Face Transformers and pipeline mechanism, loads pretrained LLM models and uses
20
+ it for inference.
21
+ """
22
+
23
+ def __init__(
24
+ self,
25
+ model_name: str,
26
+ token: Union[str, None] = None,
27
+ batch_size: int = 8,
28
+ ):
29
+ super().__init__(model_name)
30
+ self._model_name = model_name
31
+ self._token = token
32
+
33
+ self.model, self.tokenizer = hugging_face_language_model_tokenizer_factory(
34
+ model_name=self._model_name,
35
+ huggingface_token=self._token,
36
+ )
37
+
38
+ num_params = self.model.num_parameters()
39
+
40
+ # To handle max batch size for these models.
41
+ if num_params >= 70000000000: # 70B
42
+ batch_size = 2
43
+ elif num_params >= 32000000000: # 32B
44
+ batch_size = 8
45
+ elif num_params >= 27000000000: # 27B
46
+ batch_size = 16
47
+ elif "gpt-oss" in self._model_name:
48
+ batch_size = 8 # Otherwise a lot of OOM
49
+ self._batch_size_sts22 = (
50
+ 1 # For sts22, GPT-oss get OOM for batch size higher than 1.
51
+ )
52
+ self._batch_size = batch_size
53
+
54
+ def predict(self, evaluation_dataset: Dataset, task: Task) -> List:
55
+ if task.task_name == "sts22" and "gpt-oss" in self._model_name:
56
+ # For sts22, GPT-oss get OOM for batch size higher than 1.
57
+ batch_size = self._batch_size_sts22
58
+ else:
59
+ batch_size = self._batch_size
60
+
61
+ if task.task_type == TaskType.INFERENCE:
62
+ labels = task.dataset.possible_ground_truths
63
+ self.pipeline = pipeline(
64
+ task="zero-shot-classification",
65
+ model=self.model,
66
+ tokenizer=self.tokenizer,
67
+ batch_size=batch_size,
68
+ dtype="float16",
69
+ return_full_text=False,
70
+ max_new_tokens=16,
71
+ padding=True,
72
+ truncation=True,
73
+ max_length=4096,
74
+ candidate_labels=labels,
75
+ )
76
+ if len(labels) == 2:
77
+ inference_fn = self.infer_binary
78
+ else:
79
+ inference_fn = self.infer
80
+ else:
81
+ self.pipeline = pipeline(
82
+ task="text-generation",
83
+ model=self.model,
84
+ tokenizer=self.tokenizer,
85
+ batch_size=batch_size,
86
+ dtype="float16",
87
+ return_full_text=False,
88
+ max_new_tokens=64,
89
+ padding=True,
90
+ truncation=True,
91
+ max_length=4096,
92
+ )
93
+ inference_fn = self.generate
94
+
95
+ process_dataset = evaluation_dataset.map(
96
+ inference_fn,
97
+ batched=True,
98
+ batch_size=self._batch_size,
99
+ desc=f"Running evaluation for task: {task.task_name}",
100
+ remove_columns="text",
101
+ )
102
+
103
+ return list(process_dataset["prediction"])
104
+
105
+ def generate(self, rows: LazyRow) -> Dict:
106
+ """
107
+ Do a generation over a set of rows and extract the generated text and apply string post-processing.
108
+ """
109
+ with torch.no_grad():
110
+ text = rows["text"]
111
+
112
+ if self._model_name.lower() == "chocolatine":
113
+ # Problem with Phi-4 generation:
114
+ # https://github.com/huggingface/transformers/issues/36071#issuecomment-3109331152
115
+ generation_args = {"use_cache": False}
116
+ outputs = self.pipeline(text, **generation_args)
117
+ else:
118
+ outputs = self.pipeline(text)
119
+
120
+ generated_texts = [
121
+ output[0]["generated_text"].strip() for output in outputs
122
+ ]
123
+
124
+ return {"prediction": generated_texts}
125
+
126
+ def infer(self, rows: LazyRow) -> Dict:
127
+ """
128
+ Do a zero-shot classification and extract the label using a per-element generation.
129
+
130
+ For a fucking strange reason, the pipeline does not work in this case:
131
+ 1. Batched generation of more than one element
132
+ 2. More than 2 labels.
133
+
134
+ Thus, we need to loop over the element. Painful I know.
135
+ """
136
+
137
+ with torch.no_grad():
138
+ texts = rows["text"]
139
+
140
+ classifications = []
141
+ for text in texts:
142
+ output = self.pipeline(text)
143
+ classifications.append(
144
+ output["labels"][0]
145
+ ) # Labels are sorted in likelihood order.
146
+
147
+ return {"prediction": classifications}
148
+
149
+ def infer_binary(self, rows: LazyRow) -> Dict:
150
+ """
151
+ Do a binary zero-shot classification and extract the label using a per-element generation.
152
+ """
153
+ with torch.no_grad():
154
+ texts = rows["text"]
155
+
156
+ outputs = self.pipeline(texts)
157
+ classifications = [
158
+ output["labels"][0] for output in outputs
159
+ ] # Labels are sorted in likelihood order.
160
+
161
+ return {"prediction": classifications}
cole/language_model/huggingface_language_model_factory.py ADDED
@@ -0,0 +1,68 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+ from transformers import (
3
+ AutoTokenizer,
4
+ BitsAndBytesConfig,
5
+ AutoModelForCausalLM,
6
+ )
7
+ from unsloth import FastLanguageModel
8
+
9
+
10
+ def hugging_face_language_model_tokenizer_factory(
11
+ model_name,
12
+ huggingface_token: str,
13
+ ):
14
+ if (
15
+ "chocolatine" in model_name.lower()
16
+ or "lucie" in model_name.lower()
17
+ or "mixtral" in model_name.lower()
18
+ or "eurollm" in model_name.lower()
19
+ or "ibm-granite" in model_name.lower()
20
+ or "swiss-ai" in model_name.lower()
21
+ ):
22
+
23
+ compute_dtype = getattr(torch, "bfloat16")
24
+ bnb_configs = BitsAndBytesConfig(
25
+ load_in_4bit=True,
26
+ bnb_4bit_quant_type="nf4",
27
+ bnb_4bit_compute_dtype=compute_dtype,
28
+ bnb_4bit_use_double_quant=True,
29
+ )
30
+ if "chocolatine" in model_name.lower() or "mixtral" in model_name.lower():
31
+ attn_implementation = "flash_attention_2"
32
+ else:
33
+ attn_implementation = "sdpa"
34
+ model = AutoModelForCausalLM.from_pretrained(
35
+ model_name,
36
+ token=huggingface_token,
37
+ quantization_config=bnb_configs,
38
+ load_in_8bit=False, # Since we use 4bits
39
+ trust_remote_code=True,
40
+ attn_implementation=attn_implementation,
41
+ dtype=torch.float16,
42
+ )
43
+ if "chocolatine" in model_name.lower():
44
+ extra_args = {"padding_side": "left"}
45
+ else:
46
+ extra_args = {}
47
+
48
+ if "eurollm" in model_name.lower():
49
+ model.gradient_checkpointing_enable()
50
+
51
+ tokenizer = AutoTokenizer.from_pretrained(
52
+ model_name, token=huggingface_token, **extra_args
53
+ )
54
+
55
+ if tokenizer.pad_token is None:
56
+ tokenizer.pad_token_id = model.config.eos_token_id
57
+ else:
58
+ model, tokenizer = FastLanguageModel.from_pretrained(
59
+ model_name,
60
+ max_seq_length=4096,
61
+ device_map="sequential",
62
+ dtype=None,
63
+ load_in_4bit=True,
64
+ token=huggingface_token,
65
+ )
66
+
67
+ model.eval()
68
+ return model, tokenizer
cole/language_model/init_function_calling.py ADDED
@@ -0,0 +1,51 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from typing import List
2
+
3
+
4
+ def init_function_calling(labels: List[str], tool_choices: str, open_ai_api_call: bool):
5
+ if open_ai_api_call:
6
+ call = {
7
+ "tools": [
8
+ {
9
+ "type": "function",
10
+ "function": {
11
+ "name": "classification",
12
+ "description": "Use this function to return your response to the user question.",
13
+ "parameters": {
14
+ "type": "object",
15
+ "properties": {
16
+ "category": {
17
+ "type": "string",
18
+ "enum": labels,
19
+ "description": "The permitted categories to response to the question.",
20
+ },
21
+ },
22
+ "required": ["category"],
23
+ },
24
+ },
25
+ }
26
+ ],
27
+ "tool_choice": tool_choices,
28
+ }
29
+ else:
30
+ # Anthropic call
31
+ call = {
32
+ "tools": [
33
+ {
34
+ "name": "classification",
35
+ "description": "Use this function to return your response to the user question.",
36
+ "input_schema": {
37
+ "type": "object",
38
+ "properties": {
39
+ "category": {
40
+ "type": "string",
41
+ "enum": labels,
42
+ "description": "The permitted categories to response to the question.",
43
+ },
44
+ },
45
+ "required": ["category"],
46
+ },
47
+ },
48
+ ],
49
+ "tool_choice": tool_choices,
50
+ }
51
+ return call
cole/language_model/language_model_abstraction.py ADDED
@@ -0,0 +1,26 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # pylint: disable=method-hidden
2
+
3
+ from abc import abstractmethod, ABC
4
+ from typing import List
5
+
6
+ from datasets import Dataset
7
+ from datasets.formatting.formatting import LazyRow
8
+
9
+ from cole.task.task import Task
10
+
11
+
12
+ class LanguageModel(ABC):
13
+ def __init__(self, model_name: str):
14
+ self.name = model_name
15
+
16
+ @abstractmethod
17
+ def predict(self, evaluation_dataset: Dataset, task: Task) -> List:
18
+ raise NotImplementedError
19
+
20
+ @abstractmethod
21
+ def infer(self, rows: LazyRow):
22
+ raise NotImplementedError
23
+
24
+ @abstractmethod
25
+ def generate(self, rows: LazyRow):
26
+ raise NotImplementedError
cole/language_model/mistral_wrapper.py ADDED
@@ -0,0 +1,20 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from typing import Dict
2
+
3
+ from mistralai import Mistral
4
+
5
+ from cole.language_model.open_ai_api_lm_wrapper import OpenAIAPILMWrapper
6
+
7
+
8
+ class MistralWrapper(OpenAIAPILMWrapper):
9
+ def __init__(self, model_name: str, api_key: str, extra_params: Dict):
10
+ super().__init__(model_name=model_name, extra_params=extra_params)
11
+ self.client = Mistral(api_key=str(api_key))
12
+ self.tool_choices = "any"
13
+
14
+ def _inner_generate_fn(self, prompt: Dict):
15
+ return self.client.chat.complete(
16
+ model=self.model_name,
17
+ messages=prompt,
18
+ n=1,
19
+ **self._extra_params,
20
+ )
cole/language_model/open_ai_api_lm_wrapper.py ADDED
@@ -0,0 +1,171 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # pylint: disable=inconsistent-return-statements
2
+ import logging
3
+ import time
4
+ from abc import ABC, abstractmethod
5
+ from typing import List, Dict
6
+
7
+ from cole import NA_VALUE
8
+ from cole.language_model.init_function_calling import init_function_calling
9
+
10
+
11
+ class OpenAIAPILMWrapper(ABC):
12
+ def __init__(
13
+ self,
14
+ model_name: str,
15
+ extra_params: Dict,
16
+ use_function_calling: bool = True,
17
+ max_retries: int = 10,
18
+ ):
19
+ self.model_name = model_name
20
+ self._extra_params = extra_params
21
+ self.max_retries = max_retries
22
+ self.use_function_calling = use_function_calling
23
+ self.tool_choices = "required"
24
+
25
+ self._none_prediction_counter = 0
26
+ self._max_retries_counter = 0
27
+
28
+ @abstractmethod
29
+ def _inner_generate_fn(self, prompt: List) -> Dict:
30
+ pass
31
+
32
+ def init_function_calling(self, labels: List[str], tool_choices: str) -> None:
33
+ if self.use_function_calling:
34
+ open_ai_api_call = (
35
+ "claude" not in self.model_name.lower()
36
+ ) # False for Anthropic model, True for the rest
37
+ self._extra_params.update(
38
+ init_function_calling(
39
+ labels=labels,
40
+ tool_choices=tool_choices,
41
+ open_ai_api_call=open_ai_api_call,
42
+ )
43
+ )
44
+
45
+ def predict(self, text: str) -> Dict:
46
+ prompt = self.format_prompt(text)
47
+ generated_completion = self.language_model_calling(prompt=prompt)
48
+
49
+ final_prediction = self.extract_final_prediction(generated_completion)
50
+
51
+ return {"prediction": final_prediction}
52
+
53
+ @staticmethod
54
+ def format_prompt(text: str) -> List:
55
+ return [
56
+ {
57
+ "role": "user",
58
+ "content": [
59
+ {"type": "text", "text": text},
60
+ ],
61
+ }
62
+ ]
63
+
64
+ def language_model_calling(self, prompt: List):
65
+ return self._try_again(prompt=prompt)
66
+
67
+ def _try_again(self, prompt: List, retries: int = 0):
68
+ generated_completion = None
69
+ if retries > self.max_retries:
70
+ logging_message = f"Max retries exceeded: {retries}."
71
+ logging.warning(logging_message)
72
+ self._max_retries_counter += 1
73
+ return generated_completion
74
+ try:
75
+ generated_completion = self._inner_generate_fn(prompt=prompt)
76
+ return generated_completion
77
+ except:
78
+ time.sleep(5)
79
+ self._try_again(prompt=prompt, retries=retries + 1)
80
+
81
+ def extract_final_prediction(self, generated_completion) -> str:
82
+ if "claude" in self.model_name.lower():
83
+ if generated_completion is None:
84
+ final_prediction = None
85
+ elif generated_completion.content is None:
86
+ final_prediction = None
87
+ elif generated_completion.content[0].input is None:
88
+ final_prediction = None
89
+ else:
90
+ try:
91
+ final_prediction = generated_completion.content[0].input.get(
92
+ "category"
93
+ )
94
+ except:
95
+ # Case where the prediction is not a proper dictionary.
96
+ final_prediction = generated_completion.content[0].input
97
+ else:
98
+ # Something the completion is incomplete, thus we validate that components are there.
99
+ if generated_completion is None:
100
+ final_prediction = None
101
+ elif generated_completion.choices is None:
102
+ final_prediction = None
103
+ elif generated_completion.choices[0].message is None:
104
+ final_prediction = None
105
+ elif generated_completion.choices[0].message.tool_calls is None:
106
+ if generated_completion.choices[0].message.content is None:
107
+ final_prediction = None
108
+ else:
109
+ # No tools call, but potentially a response in the raw message content.
110
+ final_prediction = (
111
+ generated_completion.choices[0]
112
+ .message.content.strip()
113
+ .replace(")", "")
114
+ .strip()
115
+ )
116
+ else:
117
+ prediction = (
118
+ generated_completion.choices[0]
119
+ .message.tool_calls[0]
120
+ .function.arguments
121
+ )
122
+ try:
123
+ final_prediction = eval(prediction).get("category")
124
+ except:
125
+ # Case where the prediction is not a proper dictionary.
126
+ final_prediction = prediction
127
+
128
+ if final_prediction is None:
129
+ # Case were final prediction is None, thus we return -1 to be able to be converted
130
+ # as int if necessary (infer-case) or left as string (generate-case).
131
+ # Thus, in both case, it will not yield better results.
132
+ self._none_prediction_counter += 1
133
+ final_prediction = f"{NA_VALUE}"
134
+ elif "La réponse est" in final_prediction or ":" in final_prediction:
135
+ # To handle case where the LLM return the premise to the last query.
136
+ final_prediction = final_prediction.split(":")[-1].strip().replace(" ", "")
137
+ # Cases where the response is accompanied by other string elements, but it should be a single digit.
138
+ elif "0" in final_prediction:
139
+ final_prediction = "0"
140
+ elif "1" in final_prediction:
141
+ final_prediction = "1"
142
+ elif "2" in final_prediction:
143
+ final_prediction = "2"
144
+ elif "3" in final_prediction:
145
+ final_prediction = "3"
146
+ elif "4" in final_prediction:
147
+ final_prediction = "4"
148
+ elif "5" in final_prediction:
149
+ final_prediction = "5"
150
+ elif "6" in final_prediction:
151
+ final_prediction = "6"
152
+ elif "7" in final_prediction:
153
+ final_prediction = "7"
154
+ elif "8" in final_prediction:
155
+ final_prediction = "8"
156
+ elif "9" in final_prediction:
157
+ final_prediction = "9"
158
+ elif "10" in final_prediction:
159
+ final_prediction = "10"
160
+ elif "11" in final_prediction:
161
+ final_prediction = "11"
162
+
163
+ return final_prediction
164
+
165
+ def print_none(self) -> None:
166
+ if self._none_prediction_counter > 0:
167
+ logging_message = f"Number of None: {self._none_prediction_counter}."
168
+ logging.warning(logging_message)
169
+ if self._max_retries_counter > 0:
170
+ logging_message = f"Number of max retries exceeded occurrence: {self._max_retries_counter}."
171
+ logging.warning(logging_message)
cole/language_model/open_ai_wrapper.py ADDED
@@ -0,0 +1,32 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from typing import Dict, List, Union
2
+
3
+ from openai import OpenAI
4
+
5
+ from cole.language_model.open_ai_api_lm_wrapper import OpenAIAPILMWrapper
6
+
7
+
8
+ class OpenAIWrapper(OpenAIAPILMWrapper):
9
+ def __init__(
10
+ self,
11
+ model_name: str,
12
+ api_key: str,
13
+ extra_params: Dict,
14
+ use_function_calling: bool,
15
+ base_url: Union[None, str] = None,
16
+ timeout: int = 480,
17
+ ):
18
+ super().__init__(
19
+ model_name=model_name,
20
+ extra_params=extra_params,
21
+ use_function_calling=use_function_calling,
22
+ )
23
+ self.client = OpenAI(api_key=api_key, base_url=base_url, timeout=timeout)
24
+
25
+ def _inner_generate_fn(self, prompt: List):
26
+ return self.client.chat.completions.create(
27
+ model=self.model_name,
28
+ messages=prompt,
29
+ n=1,
30
+ stream=False,
31
+ **self._extra_params,
32
+ )
cole/language_model/private_language_model_factory.py ADDED
@@ -0,0 +1,117 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ from typing import Union
3
+
4
+ from predictions.all_llms import private_llm
5
+ from cole.language_model.anthropic_wrapper import AnthropicWrapper
6
+ from cole.language_model.cohere_wrapper import CohereWrapper
7
+ from cole.language_model.deepseek_wrapper import DeepSeekWrapper
8
+ from cole.language_model.mistral_wrapper import MistralWrapper
9
+ from cole.language_model.open_ai_wrapper import OpenAIWrapper
10
+ from cole.language_model.xai_wrapper import XAIWrapper
11
+
12
+
13
+ def get_api_key(model_name: str) -> Union[str, None]:
14
+ if model_name in private_llm["openai"]:
15
+ key_name = "openai_api_key"
16
+ elif model_name in private_llm["anthropic"]:
17
+ key_name = "anthropic_token"
18
+ elif model_name in private_llm["deepseek"]:
19
+ key_name = "deepseek_token"
20
+ elif model_name in private_llm["mistral"]:
21
+ key_name = "mistral_token"
22
+ elif model_name in private_llm["xai"]:
23
+ key_name = "XAI_API_KEY"
24
+ elif model_name in private_llm["openrouter"]:
25
+ key_name = "open_route_api_key"
26
+ elif model_name in private_llm["cohere"]:
27
+ key_name = "cohere_api_key"
28
+ else:
29
+ raise ValueError(f"Model name {model_name} not found.")
30
+
31
+ api_key = os.getenv(key_name, None)
32
+
33
+ if api_key is None:
34
+ raise ValueError(f"API key {key_name} not found.")
35
+ return api_key
36
+
37
+
38
+ def private_language_model_factory(model_name):
39
+ if model_name in private_llm["all"]:
40
+ api_key = get_api_key(model_name)
41
+
42
+ if model_name in private_llm["openai"]:
43
+ if "o1" in model_name and not "o1-mini" in model_name or "o3" in model_name:
44
+ extra_params = {
45
+ "reasoning_effort": "low"
46
+ } # Otherwise take too many tokens and stop the process.
47
+ if "o3" in model_name:
48
+ extra_params.update({"max_completion_tokens": 4096})
49
+
50
+ else:
51
+ extra_params = {}
52
+ if "mini" in model_name:
53
+ use_function_calling = False
54
+ else:
55
+ use_function_calling = True
56
+ model = OpenAIWrapper(
57
+ model_name=model_name,
58
+ api_key=api_key,
59
+ extra_params=extra_params,
60
+ use_function_calling=use_function_calling,
61
+ )
62
+ elif model_name in private_llm["anthropic"]:
63
+ extra_params = {"max_tokens": 5012}
64
+ model = AnthropicWrapper(
65
+ model_name=model_name, api_key=api_key, extra_params=extra_params
66
+ )
67
+ elif model_name in private_llm["deepseek"]:
68
+ extra_params = {"timeout": 120}
69
+ # DeepSeek reasoner does not support function calling
70
+ use_function_calling = model_name == "deepseek-reasoner"
71
+ model = DeepSeekWrapper(
72
+ model_name=model_name,
73
+ api_key=api_key,
74
+ extra_params=extra_params,
75
+ use_function_calling=use_function_calling,
76
+ )
77
+ elif model_name in private_llm["xai"]:
78
+ extra_params = {}
79
+ model = XAIWrapper(
80
+ model_name=model_name, api_key=api_key, extra_params=extra_params
81
+ )
82
+ elif model_name in private_llm["mistral"]:
83
+ extra_params = {}
84
+ model = MistralWrapper(
85
+ model_name=model_name, api_key=api_key, extra_params=extra_params
86
+ )
87
+ elif model_name in private_llm["openrouter"]:
88
+ extra_params = {}
89
+ use_function_calling = True
90
+ model = OpenAIWrapper(
91
+ model_name=model_name,
92
+ api_key=api_key,
93
+ extra_params=extra_params,
94
+ use_function_calling=use_function_calling,
95
+ base_url="https://openrouter.ai/api/v1",
96
+ timeout=480,
97
+ )
98
+ elif model_name in private_llm["cohere"]:
99
+ extra_params = {}
100
+ if "c4ai-aya" in model_name:
101
+ use_function_calling = False
102
+ else:
103
+ use_function_calling = True
104
+ model = CohereWrapper(
105
+ model_name=model_name,
106
+ api_key=api_key,
107
+ extra_params=extra_params,
108
+ use_function_calling=use_function_calling,
109
+ base_url="https://api.cohere.ai/compatibility/v1",
110
+ timeout=480,
111
+ )
112
+ else:
113
+ raise NotImplementedError("Not implemented yet.")
114
+ else:
115
+ raise ValueError(f"Model name {model_name} not found.")
116
+
117
+ return model
cole/language_model/private_lm.py ADDED
@@ -0,0 +1,64 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from typing import Union, List
2
+
3
+ from datasets import Dataset
4
+ from datasets.formatting.formatting import LazyRow
5
+
6
+ from cole.language_model.language_model_abstraction import LanguageModel
7
+ from cole.language_model.private_language_model_factory import (
8
+ private_language_model_factory,
9
+ )
10
+ from cole.task.task import Task, TaskType
11
+
12
+
13
+ class RemoteLLMModel(LanguageModel):
14
+ """
15
+ LLM Model based on private remote LLM provider (e.g. OpenAI) and pipeline mechanism for inference.
16
+ """
17
+
18
+ def generate(self, rows: LazyRow) -> Union[str, List[str]]:
19
+ try:
20
+ generate = self.model.predict(rows["text"])
21
+ except:
22
+ generate = "NA"
23
+ print("Generated a NA when model inference.")
24
+ return generate
25
+
26
+ def infer(self, rows: LazyRow) -> Union[str, List[str]]:
27
+ try:
28
+ infer = self.model.predict(rows["text"])
29
+ except:
30
+ infer = "NA"
31
+ print("Generated a NA when model inference.")
32
+ return infer
33
+
34
+ def __init__(
35
+ self,
36
+ model_name: str,
37
+ token: Union[str, None] = None,
38
+ ):
39
+ super().__init__(model_name)
40
+ self._model_name = model_name
41
+ self._token = token
42
+
43
+ self.model = private_language_model_factory(model_name=self._model_name)
44
+
45
+ def predict(self, evaluation_dataset: Dataset, task: Task) -> List:
46
+ if task.task_type == TaskType.INFERENCE:
47
+ labels = task.dataset.possible_ground_truths
48
+ self.model.init_function_calling(
49
+ labels, tool_choices=self.model.tool_choices
50
+ )
51
+ inference_fn = self.infer
52
+ else:
53
+ inference_fn = self.generate
54
+
55
+ process_dataset = evaluation_dataset.map(
56
+ inference_fn,
57
+ batched=False,
58
+ desc=f"Running evaluation for task: {task.task_name}",
59
+ remove_columns="text",
60
+ )
61
+
62
+ self.model.print_none()
63
+
64
+ return list(process_dataset["prediction"])
cole/language_model/xai_wrapper.py ADDED
@@ -0,0 +1,19 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from typing import Dict, List
2
+
3
+ from openai import OpenAI
4
+
5
+ from cole.language_model.open_ai_api_lm_wrapper import OpenAIAPILMWrapper
6
+
7
+
8
+ class XAIWrapper(OpenAIAPILMWrapper):
9
+ def __init__(self, model_name: str, api_key: str, extra_params: Dict):
10
+ super().__init__(model_name=model_name, extra_params=extra_params)
11
+ self.client = OpenAI(api_key=api_key, base_url="https://api.x.ai/v1")
12
+
13
+ def _inner_generate_fn(self, prompt: List):
14
+ return self.client.chat.completions.create(
15
+ model=self.model_name,
16
+ messages=prompt,
17
+ stream=False,
18
+ **self._extra_params,
19
+ )
cole/metrics/__init__.py ADDED
File without changes
cole/metrics/fquad_metric.py ADDED
@@ -0,0 +1,108 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import re
2
+ import string
3
+ from collections import Counter
4
+ from typing import Dict, List
5
+
6
+ from cole.metrics.metrics_wrapper import Metric
7
+
8
+
9
+ def normalize_answer(answer: str) -> str:
10
+ """
11
+ Lower text and remove punctuation, articles and extra whitespace.
12
+ Based on the SQUAD official metric: https://huggingface.co/spaces/evaluate-metric/squad
13
+ """
14
+
15
+ def remove_articles(text):
16
+ # Articles must be removed BEFORE punctuation, because the elided form
17
+ # "l'" carries its apostrophe; once `remove_punc` strips it, "l'eau" becomes
18
+ # "leau" and can no longer be matched. Order matters in French.
19
+ return re.sub(r"\b(le|la|l'|du|des|aux|un|une)\b", " ", text)
20
+
21
+ def white_space_fix(text):
22
+ return " ".join(text.split())
23
+
24
+ def remove_punc(text):
25
+ exclude = set(string.punctuation) | {"«", "»", "’", "“", "”"}
26
+ apostrophes = {"'", "’", "‘"}
27
+ # Delete apostrophes (as in the original SQuAD metric), but replace other
28
+ # punctuation with spaces to avoid gluing tokens together (e.g. in math expressions).
29
+ parts = []
30
+ for i, ch in enumerate(text):
31
+ if ch in exclude:
32
+ if ch in apostrophes:
33
+ continue
34
+ # Keep minus sign if it's followed by a digit (negative number)
35
+ if ch == "-" and i + 1 < len(text) and text[i + 1].isdigit():
36
+ parts.append(ch)
37
+ # Keep dot if it's between two digits (decimal point)
38
+ elif (
39
+ ch == "."
40
+ and i > 0
41
+ and text[i - 1].isdigit()
42
+ and i + 1 < len(text)
43
+ and text[i + 1].isdigit()
44
+ ):
45
+ parts.append(ch)
46
+ else:
47
+ parts.append(" ")
48
+ else:
49
+ parts.append(ch)
50
+ return "".join(parts)
51
+
52
+ answer = str(answer).lower()
53
+ # Normalize spaces around apostrophes (e.g., "l' eau" or "l 'eau" -> "l'eau")
54
+ answer = re.sub(r"\s*['’‘]\s*", "'", answer)
55
+ return white_space_fix(remove_punc(remove_articles(answer)))
56
+
57
+
58
+ def f1_score(prediction: str, ground_truth: str) -> float:
59
+ prediction_tokens = normalize_answer(prediction).split()
60
+ ground_truth_tokens = normalize_answer(ground_truth).split()
61
+ common = Counter(prediction_tokens) & Counter(ground_truth_tokens)
62
+ num_same = sum(common.values())
63
+ if num_same == 0:
64
+ return 0.0
65
+ precision = 1.0 * num_same / len(prediction_tokens)
66
+ recall = 1.0 * num_same / len(ground_truth_tokens)
67
+ f1 = (2 * precision * recall) / (precision + recall)
68
+ return f1
69
+
70
+
71
+ def exact_match_score(prediction: str, ground_truth: str) -> float:
72
+ return normalize_answer(prediction) == normalize_answer(ground_truth)
73
+
74
+
75
+ def metric_max_over_ground_truths(metric_fn, prediction, ground_truths):
76
+ if not ground_truths:
77
+ # No reference answer available — neither exact match nor F1 can be positive.
78
+ return 0.0
79
+ scores_for_ground_truths = []
80
+ for ground_truth in ground_truths:
81
+ score = metric_fn(prediction, ground_truth)
82
+ scores_for_ground_truths.append(score)
83
+ return max(scores_for_ground_truths)
84
+
85
+
86
+ def compute_score(predictions: List, references: List) -> Dict:
87
+ f1 = exact_match = total = 0
88
+ for prediction, reference in zip(predictions, references):
89
+ total += 1
90
+
91
+ ground_truths = reference["text"]
92
+ exact_match += metric_max_over_ground_truths(
93
+ exact_match_score, prediction, ground_truths
94
+ )
95
+ f1 += metric_max_over_ground_truths(f1_score, prediction, ground_truths)
96
+
97
+ if total == 0:
98
+ return {"exact_match": 0.0, "f1": 0.0}
99
+ exact_match = 100.0 * exact_match / total
100
+ f1 = 100.0 * f1 / total
101
+
102
+ return {"exact_match": exact_match, "f1": f1}
103
+
104
+
105
+ class FQuAD(Metric):
106
+ def compute(self, predictions: List, references: List) -> Dict:
107
+ score = compute_score(predictions=predictions, references=references)
108
+ return score