Spaces:
Running
Running
COLE CI commited on
Commit ·
424c5d9
0
Parent(s):
deploy to HF Space
Browse filesThis view is limited to 50 files because it contains too many changes. See raw diff
- .gitattributes +2 -0
- .gitignore +69 -0
- .idea/.gitignore +8 -0
- .pre-commit-config.yaml +27 -0
- CITATION.cff +40 -0
- CLAUDE.md +70 -0
- CODE_OF_CONDUCT.md +35 -0
- Dockerfile +55 -0
- LICENSE +21 -0
- Makefile +19 -0
- README.md +128 -0
- SECURITY.md +15 -0
- cole/__init__.py +6 -0
- cole/backend/__init__.py +0 -0
- cole/backend/evaluation.py +41 -0
- cole/backend/results/leaderboard.json +0 -0
- cole/backend/submission_api.py +251 -0
- cole/backend/submit_tools.py +37 -0
- cole/backend/validation_tools.py +90 -0
- cole/dataset/__init__.py +0 -0
- cole/dataset/dataset.py +101 -0
- cole/dataset/datasets_data.py +797 -0
- cole/dataset/prompt_builder.py +43 -0
- cole/docker_requirements.txt +18 -0
- cole/evaluation/__init__.py +0 -0
- cole/evaluation/evaluation_pipeline.py +159 -0
- cole/evaluation/evaluation_pipeline_private_llm.py +137 -0
- cole/evaluation/evaluation_pipeline_small.py +150 -0
- cole/evaluation/evaluation_pipeline_small_2.py +151 -0
- cole/evaluation/llm_evaluator.py +217 -0
- cole/evaluation/llm_factory.py +25 -0
- cole/evaluation/tools.py +33 -0
- cole/language_model/__init__.py +0 -0
- cole/language_model/anthropic_wrapper.py +20 -0
- cole/language_model/baseline.py +45 -0
- cole/language_model/cohere_wrapper.py +31 -0
- cole/language_model/deepseek_wrapper.py +30 -0
- cole/language_model/google_wrapper.py +22 -0
- cole/language_model/hugging_face_lm.py +161 -0
- cole/language_model/huggingface_language_model_factory.py +68 -0
- cole/language_model/init_function_calling.py +51 -0
- cole/language_model/language_model_abstraction.py +26 -0
- cole/language_model/mistral_wrapper.py +20 -0
- cole/language_model/open_ai_api_lm_wrapper.py +171 -0
- cole/language_model/open_ai_wrapper.py +32 -0
- cole/language_model/private_language_model_factory.py +117 -0
- cole/language_model/private_lm.py +64 -0
- cole/language_model/xai_wrapper.py +19 -0
- cole/metrics/__init__.py +0 -0
- 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 |
+
[](https://colebenchmark.org/)
|
| 13 |
+
[](https://arxiv.org/abs/2510.05046)
|
| 14 |
+
[](https://huggingface.co/datasets/graalul/COLE-public)
|
| 15 |
+
[](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
|