Spaces:
Running
Running
Test datasets
Browse files- README.md +18 -0
- data/training/manifest.json +4 -4
- data/training/places_retrieval_v1/README.md +27 -5
- data/training/places_retrieval_v1/test.jsonl +0 -0
- data/training/places_retrieval_v1/validation.jsonl +0 -0
- scripts/build_place_training_datasets.py +220 -8
- scripts/evaluate_place_retriever.py +250 -0
- tests/test_evaluate_place_retriever.py +51 -0
- tests/test_place_training_datasets.py +57 -11
README.md
CHANGED
|
@@ -388,6 +388,24 @@ Para fine-tuning, `scripts/train_place_retriever.py` acepta JSONL con `query`,
|
|
| 388 |
`positive` y `hard_negatives`. El artefacto resultante se configura mediante
|
| 389 |
`PLACES_EMBEDDING_MODEL`; no se incluye un modelo ficticio preentrenado en el repo.
|
| 390 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 391 |
El extractor de intencion se entrena por separado con
|
| 392 |
`scripts/train_place_intent_bert.py`. Su JSONL contiene `text` y spans abiertos
|
| 393 |
`{start, end, slot}`; los slots permitidos son `CATEGORY`, `PREFERENCE`,
|
|
|
|
| 388 |
`positive` y `hard_negatives`. El artefacto resultante se configura mediante
|
| 389 |
`PLACES_EMBEDDING_MODEL`; no se incluye un modelo ficticio preentrenado en el repo.
|
| 390 |
|
| 391 |
+
La evaluación del retriever debe usar el corpus global, no cuatro candidatos aislados
|
| 392 |
+
por fila. Validation incluye 60 consultas y el test sintético 80 consultas, cada uno
|
| 393 |
+
sobre 20 documentos propios y no vistos, con lenguaje
|
| 394 |
+
coloquial, errores ortográficos, necesidades implícitas y frases contrastivas. Para
|
| 395 |
+
comparar el modelo base y el fine-tuned sobre exactamente el mismo corpus:
|
| 396 |
+
|
| 397 |
+
```powershell
|
| 398 |
+
python scripts/evaluate_place_retriever.py `
|
| 399 |
+
--test-file data/training/places_retrieval_v1/test.jsonl `
|
| 400 |
+
--model base=intfloat/multilingual-e5-base `
|
| 401 |
+
--model fine_tuned=C:\ruta\al\places-e5-retriever-v1 `
|
| 402 |
+
--output-json artifacts/retriever-evaluation.json
|
| 403 |
+
```
|
| 404 |
+
|
| 405 |
+
Además de Top-1, Recall@k, MRR y nDCG@10, el reporte separa resultados para
|
| 406 |
+
`colloquial`, `misspelling`, `implicit`, `contrastive` y otras dificultades. El
|
| 407 |
+
prefijo `query:`/`passage:` se aplica dentro del evaluador.
|
| 408 |
+
|
| 409 |
El extractor de intencion se entrena por separado con
|
| 410 |
`scripts/train_place_intent_bert.py`. Su JSONL contiene `text` y spans abiertos
|
| 411 |
`{start, end, slot}`; los slots permitidos son `CATEGORY`, `PREFERENCE`,
|
data/training/manifest.json
CHANGED
|
@@ -27,13 +27,13 @@
|
|
| 27 |
},
|
| 28 |
"validation": {
|
| 29 |
"path": "data/training/places_retrieval_v1/validation.jsonl",
|
| 30 |
-
"records":
|
| 31 |
-
"sha256": "
|
| 32 |
},
|
| 33 |
"test": {
|
| 34 |
"path": "data/training/places_retrieval_v1/test.jsonl",
|
| 35 |
-
"records":
|
| 36 |
-
"sha256": "
|
| 37 |
}
|
| 38 |
}
|
| 39 |
}
|
|
|
|
| 27 |
},
|
| 28 |
"validation": {
|
| 29 |
"path": "data/training/places_retrieval_v1/validation.jsonl",
|
| 30 |
+
"records": 60,
|
| 31 |
+
"sha256": "9f61e3467e907f195f14015f15acf78f7d4568fb2198b86516724ce1088a8241"
|
| 32 |
},
|
| 33 |
"test": {
|
| 34 |
"path": "data/training/places_retrieval_v1/test.jsonl",
|
| 35 |
+
"records": 80,
|
| 36 |
+
"sha256": "d99fd7094e7cc999b62ee264d88280e62e464c9a9562be8ebdeb1a9f06d64b37"
|
| 37 |
}
|
| 38 |
}
|
| 39 |
}
|
data/training/places_retrieval_v1/README.md
CHANGED
|
@@ -4,7 +4,10 @@ Dataset JSONL sintético para arrancar el bi-encoder E5 de Places. Cada fila con
|
|
| 4 |
|
| 5 |
- `query`: consulta sin prefijo E5.
|
| 6 |
- `positive`: documento con formato `structured-place-v3`.
|
| 7 |
-
- `hard_negatives`:
|
|
|
|
|
|
|
|
|
|
| 8 |
|
| 9 |
Los splits no comparten queries ni documentos positivos. No agregues `query:` o
|
| 10 |
`passage:` a los archivos: el script de entrenamiento aplica esos prefijos.
|
|
@@ -12,15 +15,34 @@ Los splits no comparten queries ni documentos positivos. No agregues `query:` o
|
|
| 12 |
Conteos:
|
| 13 |
|
| 14 |
- `train.jsonl`: 160 pares.
|
| 15 |
-
- `validation.jsonl`:
|
| 16 |
-
|
| 17 |
-
-
|
|
|
|
|
|
|
| 18 |
|
| 19 |
Las sedes derivadas tienen nombres, documentos, consultas y atributos discriminantes
|
| 20 |
propios. Variantes del mismo concepto se excluyen de los hard negatives explícitos
|
| 21 |
para no etiquetar como incorrecto un lugar que también podría satisfacer la consulta.
|
| 22 |
Validation y test conservan un solo documento por cada concepto no visto; así sus positivos
|
| 23 |
-
no compiten contra una sede hermana igualmente relevante durante la evaluación.
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 24 |
|
| 25 |
Este corpus cubre vocabulario abierto y sirve para iniciar el fine-tuning, pero no es
|
| 26 |
suficiente por sí solo para una decisión productiva. Sustituir o ampliar train con
|
|
|
|
| 4 |
|
| 5 |
- `query`: consulta sin prefijo E5.
|
| 6 |
- `positive`: documento con formato `structured-place-v3`.
|
| 7 |
+
- `hard_negatives`: lugares plausibles pero incorrectos del mismo split. Train
|
| 8 |
+
contiene tres; validation y test contienen siete para evitar evaluaciones triviales.
|
| 9 |
+
- `challenge_tags` (validation y test): tipo de dificultad de la consulta, por ejemplo
|
| 10 |
+
`colloquial`, `misspelling`, `implicit`, `contrastive` o `constraint`.
|
| 11 |
|
| 12 |
Los splits no comparten queries ni documentos positivos. No agregues `query:` o
|
| 13 |
`passage:` a los archivos: el script de entrenamiento aplica esos prefijos.
|
|
|
|
| 15 |
Conteos:
|
| 16 |
|
| 17 |
- `train.jsonl`: 160 pares.
|
| 18 |
+
- `validation.jsonl`: 60 consultas sobre 20 conceptos no vistos (tres formulaciones
|
| 19 |
+
por documento relevante).
|
| 20 |
+
- `test.jsonl`: 80 consultas sobre 20 conceptos no vistos (cuatro formulaciones
|
| 21 |
+
por documento relevante).
|
| 22 |
+
- Total: 300 registros; el split usado para entrenar permanece en 160 pares.
|
| 23 |
|
| 24 |
Las sedes derivadas tienen nombres, documentos, consultas y atributos discriminantes
|
| 25 |
propios. Variantes del mismo concepto se excluyen de los hard negatives explícitos
|
| 26 |
para no etiquetar como incorrecto un lugar que también podría satisfacer la consulta.
|
| 27 |
Validation y test conservan un solo documento por cada concepto no visto; así sus positivos
|
| 28 |
+
no compiten contra una sede hermana igualmente relevante durante la evaluación. Test repite
|
| 29 |
+
el documento relevante para consultas canónicas, coloquiales, implícitas y con errores de
|
| 30 |
+
escritura. Eso es intencional: varios intents lingüísticos pueden compartir el mismo qrel.
|
| 31 |
+
Validation también contiene reformulaciones difíciles para elegir épocas y checkpoints sin
|
| 32 |
+
consultar repetidamente el holdout final.
|
| 33 |
+
|
| 34 |
+
No evalúes cada fila únicamente contra sus `hard_negatives`. Usa el corpus global
|
| 35 |
+
deduplicado para que las 80 consultas compitan contra los 20 documentos de test:
|
| 36 |
+
|
| 37 |
+
```bash
|
| 38 |
+
python scripts/evaluate_place_retriever.py \
|
| 39 |
+
--test-file data/training/places_retrieval_v1/test.jsonl \
|
| 40 |
+
--model base=intfloat/multilingual-e5-base \
|
| 41 |
+
--model fine_tuned=/ruta/al/modelo
|
| 42 |
+
```
|
| 43 |
+
|
| 44 |
+
El resultado incluye Top-1, Recall@3/5/10, MRR, nDCG@10 y métricas separadas por
|
| 45 |
+
`challenge_tags`. Elige épocas con validation; usa test solo para la comparación final.
|
| 46 |
|
| 47 |
Este corpus cubre vocabulario abierto y sirve para iniciar el fine-tuning, pero no es
|
| 48 |
suficiente por sí solo para una decisión productiva. Sustituir o ampliar train con
|
data/training/places_retrieval_v1/test.jsonl
CHANGED
|
The diff for this file is too large to render.
See raw diff
|
|
|
data/training/places_retrieval_v1/validation.jsonl
CHANGED
|
The diff for this file is too large to render.
See raw diff
|
|
|
scripts/build_place_training_datasets.py
CHANGED
|
@@ -294,6 +294,207 @@ PLACE_SPECS: tuple[PlaceSpec, ...] = (
|
|
| 294 |
)
|
| 295 |
|
| 296 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 297 |
def main() -> None:
|
| 298 |
INTENT_ROOT.mkdir(parents=True, exist_ok=True)
|
| 299 |
RETRIEVAL_ROOT.mkdir(parents=True, exist_ok=True)
|
|
@@ -409,16 +610,24 @@ def _retrieval_records(split: str) -> list[dict[str, object]]:
|
|
| 409 |
and _concept_key(candidate) != _concept_key(spec)
|
| 410 |
]
|
| 411 |
rotated = other_family[index % len(other_family) :] + other_family[: index % len(other_family)]
|
| 412 |
-
|
| 413 |
-
|
|
|
|
| 414 |
raise RuntimeError(f"Not enough hard negatives for {spec.name}")
|
| 415 |
-
|
| 416 |
-
|
| 417 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 418 |
"positive": documents[spec.name],
|
| 419 |
"hard_negatives": [documents[item.name] for item in negative_specs],
|
| 420 |
}
|
| 421 |
-
|
|
|
|
|
|
|
| 422 |
return records
|
| 423 |
|
| 424 |
|
|
@@ -544,8 +753,11 @@ def _validate_cross_split_uniqueness() -> None:
|
|
| 544 |
positive = str(record["positive"])
|
| 545 |
if query in retrieval_seen:
|
| 546 |
raise ValueError(f"Duplicate retrieval query in {retrieval_seen[query]} and {split}: {query}")
|
| 547 |
-
|
| 548 |
-
|
|
|
|
|
|
|
|
|
|
| 549 |
retrieval_seen[query] = split
|
| 550 |
positive_seen[positive] = split
|
| 551 |
|
|
|
|
| 294 |
)
|
| 295 |
|
| 296 |
|
| 297 |
+
# Validation and test deliberately contain several queries for the same relevant
|
| 298 |
+
# document. Validation supports epoch/model selection without inspecting test.
|
| 299 |
+
# Test remains the final, harder holdout.
|
| 300 |
+
# This is closer to production retrieval than a one-query/one-document lookup:
|
| 301 |
+
# the model must survive slang, misspellings, implicit needs and contrastive
|
| 302 |
+
# wording while ranking against the complete test corpus. Keep these variants
|
| 303 |
+
# out of train to prevent evaluation leakage.
|
| 304 |
+
RETRIEVAL_VALIDATION_QUERY_VARIANTS: dict[
|
| 305 |
+
str,
|
| 306 |
+
tuple[tuple[str, tuple[str, ...]], ...],
|
| 307 |
+
] = {
|
| 308 |
+
"Pupusas Centroamérica": (
|
| 309 |
+
("se me antoja comida salvadoreña de masa rellena", ("implicit", "paraphrase")),
|
| 310 |
+
("unas pupusitas con curtido dónde", ("colloquial", "implicit")),
|
| 311 |
+
),
|
| 312 |
+
"Crêpe Maison": (
|
| 313 |
+
("quiero algo enrollado que pueda ser dulce o salado", ("implicit", "contrastive")),
|
| 314 |
+
("crepas pa cenar, no solo postre", ("colloquial", "contrastive")),
|
| 315 |
+
),
|
| 316 |
+
"Planetario Orión": (
|
| 317 |
+
("quiero ver el cielo proyectado sin salir de la ciudad", ("implicit", "paraphrase")),
|
| 318 |
+
("plan de estrellas y astronomía bajo techo", ("implicit", "constraint")),
|
| 319 |
+
),
|
| 320 |
+
"Aqua Centro": (
|
| 321 |
+
("donde hay carriles pa echarme unos largos", ("colloquial", "implicit")),
|
| 322 |
+
("alberca techada con clases para nadar", ("paraphrase", "constraint")),
|
| 323 |
+
),
|
| 324 |
+
"Clave Oculta": (
|
| 325 |
+
("plan con acertijos pa escapar con la banda", ("colloquial", "implicit")),
|
| 326 |
+
("cuarto temático donde resolvemos pistas", ("implicit", "paraphrase")),
|
| 327 |
+
),
|
| 328 |
+
"Vinilo Sur": (
|
| 329 |
+
("dónde consigo discos de los grandotes", ("colloquial", "implicit")),
|
| 330 |
+
("ando cazando elepés usados", ("colloquial", "paraphrase")),
|
| 331 |
+
),
|
| 332 |
+
"Centro Presente": (
|
| 333 |
+
("necesito respirar y apagar la mente un rato", ("implicit", "colloquial")),
|
| 334 |
+
("un lugar silencioso con meditación guiada", ("paraphrase", "constraint")),
|
| 335 |
+
),
|
| 336 |
+
"Lúpulo Local": (
|
| 337 |
+
("cheve artesanal y que expliquen cómo la hacen", ("colloquial", "implicit")),
|
| 338 |
+
("quiero degustar varias cervezas locales", ("paraphrase",)),
|
| 339 |
+
),
|
| 340 |
+
"Sushi Viajero": (
|
| 341 |
+
("sushi que vaya pasando frente a la mesa", ("implicit", "paraphrase")),
|
| 342 |
+
("japonés con banda de platitos", ("colloquial", "implicit")),
|
| 343 |
+
),
|
| 344 |
+
"Cero Residuo": (
|
| 345 |
+
("comprar despensa llevando mis propios frascos", ("implicit", "constraint")),
|
| 346 |
+
("tienda a granel sin tanto empaque", ("paraphrase", "constraint")),
|
| 347 |
+
),
|
| 348 |
+
"Salto Alto": (
|
| 349 |
+
("dónde llevo a brincar a los chamacos", ("colloquial", "implicit")),
|
| 350 |
+
("parque techado lleno de camas elásticas", ("paraphrase", "implicit")),
|
| 351 |
+
),
|
| 352 |
+
"Teatro Bel Canto": (
|
| 353 |
+
("quiero escuchar voces líricas en un teatro", ("implicit", "paraphrase")),
|
| 354 |
+
("algún recinto con bel canto y música de cámara", ("paraphrase",)),
|
| 355 |
+
),
|
| 356 |
+
"Finca Aroma": (
|
| 357 |
+
("quiero ver de dónde sale el café antes de la taza", ("implicit", "paraphrase")),
|
| 358 |
+
("paseo entre cafetales con tostado", ("implicit", "paraphrase")),
|
| 359 |
+
),
|
| 360 |
+
"Raqueta Sur": (
|
| 361 |
+
("dónde rentan raqueta y cancha de squash", ("constraint", "paraphrase")),
|
| 362 |
+
("quiero pegarle a la pelota contra la pared", ("implicit", "colloquial")),
|
| 363 |
+
),
|
| 364 |
+
"Fábrica Lab": (
|
| 365 |
+
("taller para fabricar prototipos con impresora 3d", ("implicit", "paraphrase")),
|
| 366 |
+
("necesito herramientas y electrónica compartidas", ("implicit", "constraint")),
|
| 367 |
+
),
|
| 368 |
+
"Estudio Obturador": (
|
| 369 |
+
("renta de ciclorama y luces por hora", ("constraint", "paraphrase")),
|
| 370 |
+
("lugar equipado pa una sesión de fotos", ("colloquial", "implicit")),
|
| 371 |
+
),
|
| 372 |
+
"Verso Café": (
|
| 373 |
+
("cafecito con micro abierto para leer versos", ("colloquial", "implicit")),
|
| 374 |
+
("dónde escucho poesía mientras tomo algo", ("implicit", "paraphrase")),
|
| 375 |
+
),
|
| 376 |
+
"Baño Nube": (
|
| 377 |
+
("quiero un baño de vapor árabe", ("implicit", "paraphrase")),
|
| 378 |
+
("spa con circuito de vapor tipo turco", ("paraphrase", "constraint")),
|
| 379 |
+
),
|
| 380 |
+
"Cosecha Directa": (
|
| 381 |
+
("verduras directo de quien las cosecha", ("implicit", "paraphrase")),
|
| 382 |
+
("mercadito semanal de productores", ("colloquial", "paraphrase")),
|
| 383 |
+
),
|
| 384 |
+
"Taco Verde": (
|
| 385 |
+
("tacos sin nada de origen animal", ("implicit", "constraint")),
|
| 386 |
+
("taquitos de setas, cero carne", ("colloquial", "contrastive")),
|
| 387 |
+
),
|
| 388 |
+
}
|
| 389 |
+
|
| 390 |
+
|
| 391 |
+
RETRIEVAL_TEST_QUERY_VARIANTS: dict[
|
| 392 |
+
str,
|
| 393 |
+
tuple[tuple[str, tuple[str, ...]], ...],
|
| 394 |
+
] = {
|
| 395 |
+
"Arepa Morena": (
|
| 396 |
+
("traigo antojo de una reina pepiada", ("colloquial", "implicit")),
|
| 397 |
+
("onde venden arepitas chidas", ("colloquial", "misspelling")),
|
| 398 |
+
("algo venezolano relleno y hecho al momento", ("implicit", "paraphrase")),
|
| 399 |
+
),
|
| 400 |
+
"Michi Café": (
|
| 401 |
+
("un cafecito pa echar el rato con michis", ("colloquial", "implicit")),
|
| 402 |
+
("cafeteria con gatitos rescatados", ("misspelling", "paraphrase")),
|
| 403 |
+
("quiero tomar algo rodeado de mininos", ("implicit", "paraphrase")),
|
| 404 |
+
),
|
| 405 |
+
"Laboratorio Curioso": (
|
| 406 |
+
("algo chido pa que los peques hagan experimentos", ("colloquial", "implicit")),
|
| 407 |
+
("museo no, mejor talleres cientificos para morritos", ("contrastive", "colloquial")),
|
| 408 |
+
("donde llevo a mi hijo a aprender ciencia jugando", ("implicit", "paraphrase")),
|
| 409 |
+
),
|
| 410 |
+
"Santuario Monarca": (
|
| 411 |
+
("quiero ver un buen de mariposas en un jardin", ("colloquial", "misspelling")),
|
| 412 |
+
("algún lugar de conservación de monarcas", ("implicit", "paraphrase")),
|
| 413 |
+
("plan tranqui entre plantas y maripositas", ("colloquial", "implicit")),
|
| 414 |
+
),
|
| 415 |
+
"Risa Abierta": (
|
| 416 |
+
("donde hay standap esta noche", ("misspelling", "colloquial")),
|
| 417 |
+
("quiero echarme unas risas con micro abierto", ("colloquial", "implicit")),
|
| 418 |
+
("un plan de comedia en vivo, no cine", ("contrastive", "paraphrase")),
|
| 419 |
+
),
|
| 420 |
+
"Café Babel": (
|
| 421 |
+
("cafecito pa practicar inglés con la banda", ("colloquial", "implicit")),
|
| 422 |
+
("intercambio d idiomas en un lugar para tomar algo", ("misspelling", "paraphrase")),
|
| 423 |
+
("quiero conversar con gente en otras lenguas", ("implicit", "paraphrase")),
|
| 424 |
+
),
|
| 425 |
+
"Miga Verde": (
|
| 426 |
+
("algo dulce pero cero ingredientes animales", ("implicit", "contrastive")),
|
| 427 |
+
("panecito vegano pa'l antojo", ("colloquial", "paraphrase")),
|
| 428 |
+
("pasteles sin leche ni huevo, qué hay", ("implicit", "constraint")),
|
| 429 |
+
),
|
| 430 |
+
"Bosque Domo": (
|
| 431 |
+
("acampar sí pero sin sufrirle", ("colloquial", "implicit")),
|
| 432 |
+
("un domito cómodo en el bosque", ("colloquial", "paraphrase")),
|
| 433 |
+
("naturaleza con cama y techo, plan de finde", ("implicit", "colloquial")),
|
| 434 |
+
),
|
| 435 |
+
"La Media Luna": (
|
| 436 |
+
("se me antoja algo relleno y horneadito", ("implicit", "colloquial")),
|
| 437 |
+
("onde hay empanadas dulces y saladas", ("misspelling", "paraphrase")),
|
| 438 |
+
("un lugar de medias lunas no, de empanadas", ("contrastive", "adversarial")),
|
| 439 |
+
),
|
| 440 |
+
"Herpetario Verde": (
|
| 441 |
+
("quiero ver víboras y lagartos", ("implicit", "paraphrase")),
|
| 442 |
+
("un plan educativo con animalitos de sangre fría", ("implicit", "paraphrase")),
|
| 443 |
+
("donde conocen los peques a los reptiles", ("colloquial", "implicit")),
|
| 444 |
+
),
|
| 445 |
+
"Museo del Barro": (
|
| 446 |
+
("quiero ver piezas antiguas hechas de barro", ("implicit", "paraphrase")),
|
| 447 |
+
("museo d alfareria de la region", ("misspelling", "paraphrase")),
|
| 448 |
+
("un plan cultural sobre vasijas y cerámica", ("implicit", "paraphrase")),
|
| 449 |
+
),
|
| 450 |
+
"Pista Furia": (
|
| 451 |
+
("dónde entrenan ese deporte rudo en patines", ("implicit", "paraphrase")),
|
| 452 |
+
("quiero darle al roler derbi", ("misspelling", "colloquial")),
|
| 453 |
+
("pista para aprender derby sobre ruedas", ("paraphrase", "implicit")),
|
| 454 |
+
),
|
| 455 |
+
"Voz Estudio": (
|
| 456 |
+
("necesito una cabina que no meta ruido pa grabar", ("colloquial", "implicit")),
|
| 457 |
+
("donde puedo producir mi podcast", ("paraphrase",)),
|
| 458 |
+
("un estudio con tratamiento acustico y edición", ("misspelling", "implicit")),
|
| 459 |
+
),
|
| 460 |
+
"Invernadero Café": (
|
| 461 |
+
("cafecito entre un montón de plantas", ("colloquial", "implicit")),
|
| 462 |
+
("un lugar verde para tomar café sin estar afuera", ("contrastive", "implicit")),
|
| 463 |
+
("cafeteria tipo jardin techado", ("misspelling", "paraphrase")),
|
| 464 |
+
),
|
| 465 |
+
"Robot Peques": (
|
| 466 |
+
("algo pa que mi morrito arme robots", ("colloquial", "implicit")),
|
| 467 |
+
("tayer infantil de programacion y electronica", ("misspelling", "paraphrase")),
|
| 468 |
+
("donde enseñan tecnología construyendo cosas", ("implicit", "paraphrase")),
|
| 469 |
+
),
|
| 470 |
+
"Remo Río": (
|
| 471 |
+
("quiero remar pero necesito que me presten todo", ("implicit", "constraint")),
|
| 472 |
+
("renta de kayak pa novatos con alguien que guíe", ("colloquial", "constraint")),
|
| 473 |
+
("plan en el río con remo y equipo seguro", ("implicit", "paraphrase")),
|
| 474 |
+
),
|
| 475 |
+
"Quesería Sierra": (
|
| 476 |
+
("quiero probar quesitos de la región", ("colloquial", "implicit")),
|
| 477 |
+
("una tienda donde den degustacion de queso artesanal", ("misspelling", "paraphrase")),
|
| 478 |
+
("ando buscando lácteos locales, sobre todo quesos", ("implicit", "paraphrase")),
|
| 479 |
+
),
|
| 480 |
+
"Circo Aire": (
|
| 481 |
+
("quiero aprender a colgarme de telas sin matarme", ("colloquial", "implicit")),
|
| 482 |
+
("clases d acrobacia aerea pa principiantes", ("misspelling", "colloquial")),
|
| 483 |
+
("una escuela de circo con equilibrio y telas", ("paraphrase",)),
|
| 484 |
+
),
|
| 485 |
+
"Reserva Alas": (
|
| 486 |
+
("plan para pajarear con binoculares", ("colloquial", "implicit")),
|
| 487 |
+
("un sitio con observatorios para avistar plumíferos", ("paraphrase", "implicit")),
|
| 488 |
+
("quiero caminar por una reserva viendo aves", ("paraphrase",)),
|
| 489 |
+
),
|
| 490 |
+
"Bar Cero": (
|
| 491 |
+
("quiero salir de noche sin ponerme peda", ("colloquial", "implicit")),
|
| 492 |
+
("bar con tragos cero alcohol y ambiente tranqui", ("colloquial", "constraint")),
|
| 493 |
+
("mocktels y musiquita pero nada de chela", ("misspelling", "colloquial", "constraint")),
|
| 494 |
+
),
|
| 495 |
+
}
|
| 496 |
+
|
| 497 |
+
|
| 498 |
def main() -> None:
|
| 499 |
INTENT_ROOT.mkdir(parents=True, exist_ok=True)
|
| 500 |
RETRIEVAL_ROOT.mkdir(parents=True, exist_ok=True)
|
|
|
|
| 610 |
and _concept_key(candidate) != _concept_key(spec)
|
| 611 |
]
|
| 612 |
rotated = other_family[index % len(other_family) :] + other_family[: index % len(other_family)]
|
| 613 |
+
negative_limit = 3 if split == "train" else 7
|
| 614 |
+
negative_specs = _unique_specs((*same_family, *rotated))[:negative_limit]
|
| 615 |
+
if len(negative_specs) < negative_limit:
|
| 616 |
raise RuntimeError(f"Not enough hard negatives for {spec.name}")
|
| 617 |
+
query_variants = [(spec.query, ("canonical",))]
|
| 618 |
+
if split == "validation":
|
| 619 |
+
query_variants.extend(RETRIEVAL_VALIDATION_QUERY_VARIANTS[spec.name])
|
| 620 |
+
elif split == "test":
|
| 621 |
+
query_variants.extend(RETRIEVAL_TEST_QUERY_VARIANTS[spec.name])
|
| 622 |
+
for query, challenge_tags in query_variants:
|
| 623 |
+
record: dict[str, object] = {
|
| 624 |
+
"query": query,
|
| 625 |
"positive": documents[spec.name],
|
| 626 |
"hard_negatives": [documents[item.name] for item in negative_specs],
|
| 627 |
}
|
| 628 |
+
if split != "train":
|
| 629 |
+
record["challenge_tags"] = list(challenge_tags)
|
| 630 |
+
records.append(record)
|
| 631 |
return records
|
| 632 |
|
| 633 |
|
|
|
|
| 753 |
positive = str(record["positive"])
|
| 754 |
if query in retrieval_seen:
|
| 755 |
raise ValueError(f"Duplicate retrieval query in {retrieval_seen[query]} and {split}: {query}")
|
| 756 |
+
positive_owner = positive_seen.get(positive)
|
| 757 |
+
if positive_owner is not None and positive_owner != split:
|
| 758 |
+
raise ValueError(
|
| 759 |
+
f"Duplicate positive document in {positive_owner} and {split}"
|
| 760 |
+
)
|
| 761 |
retrieval_seen[query] = split
|
| 762 |
positive_seen[positive] = split
|
| 763 |
|
scripts/evaluate_place_retriever.py
ADDED
|
@@ -0,0 +1,250 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Evaluate E5 place retrievers against one shared, global candidate corpus.
|
| 2 |
+
|
| 3 |
+
Unlike a per-row positive-versus-negatives check, every query is ranked against
|
| 4 |
+
all unique positives and hard negatives in the JSONL file. This avoids the
|
| 5 |
+
four-candidate ceiling that made both the base and fine-tuned models score 1.0.
|
| 6 |
+
"""
|
| 7 |
+
|
| 8 |
+
from __future__ import annotations
|
| 9 |
+
|
| 10 |
+
import argparse
|
| 11 |
+
import json
|
| 12 |
+
from collections import defaultdict
|
| 13 |
+
from pathlib import Path
|
| 14 |
+
from typing import Any, Iterable
|
| 15 |
+
|
| 16 |
+
import numpy as np
|
| 17 |
+
|
| 18 |
+
|
| 19 |
+
DEFAULT_BASE_MODEL = "intfloat/multilingual-e5-base"
|
| 20 |
+
|
| 21 |
+
|
| 22 |
+
def read_evaluation_rows(path: Path) -> list[dict[str, Any]]:
|
| 23 |
+
if not path.is_file():
|
| 24 |
+
raise FileNotFoundError(f"Evaluation dataset not found: {path}")
|
| 25 |
+
|
| 26 |
+
rows: list[dict[str, Any]] = []
|
| 27 |
+
for line_number, raw_line in enumerate(
|
| 28 |
+
path.read_text(encoding="utf-8").splitlines(),
|
| 29 |
+
start=1,
|
| 30 |
+
):
|
| 31 |
+
if not raw_line.strip():
|
| 32 |
+
continue
|
| 33 |
+
payload = json.loads(raw_line)
|
| 34 |
+
query = _required_text(payload.get("query"), line_number, "query")
|
| 35 |
+
positive = _required_text(payload.get("positive"), line_number, "positive")
|
| 36 |
+
negatives = payload.get("hard_negatives", [])
|
| 37 |
+
if not isinstance(negatives, list):
|
| 38 |
+
raise ValueError(
|
| 39 |
+
f"Line {line_number}: hard_negatives must be a list of strings"
|
| 40 |
+
)
|
| 41 |
+
hard_negatives = [
|
| 42 |
+
_required_text(value, line_number, "hard_negatives")
|
| 43 |
+
for value in negatives
|
| 44 |
+
]
|
| 45 |
+
raw_tags = payload.get("challenge_tags", ["unclassified"])
|
| 46 |
+
if not isinstance(raw_tags, list) or not raw_tags:
|
| 47 |
+
raise ValueError(
|
| 48 |
+
f"Line {line_number}: challenge_tags must be a non-empty list"
|
| 49 |
+
)
|
| 50 |
+
challenge_tags = [
|
| 51 |
+
_required_text(value, line_number, "challenge_tags")
|
| 52 |
+
for value in raw_tags
|
| 53 |
+
]
|
| 54 |
+
rows.append(
|
| 55 |
+
{
|
| 56 |
+
"query": query,
|
| 57 |
+
"positive": positive,
|
| 58 |
+
"hard_negatives": hard_negatives,
|
| 59 |
+
"challenge_tags": challenge_tags,
|
| 60 |
+
}
|
| 61 |
+
)
|
| 62 |
+
|
| 63 |
+
if not rows:
|
| 64 |
+
raise ValueError("At least one evaluation example is required")
|
| 65 |
+
return rows
|
| 66 |
+
|
| 67 |
+
|
| 68 |
+
def build_global_corpus(rows: Iterable[dict[str, Any]]) -> list[str]:
|
| 69 |
+
"""Return all candidate documents once, preserving their first-seen order."""
|
| 70 |
+
|
| 71 |
+
corpus: list[str] = []
|
| 72 |
+
seen: set[str] = set()
|
| 73 |
+
materialized = list(rows)
|
| 74 |
+
for row in materialized:
|
| 75 |
+
candidates = (row["positive"], *row.get("hard_negatives", []))
|
| 76 |
+
for candidate in candidates:
|
| 77 |
+
if candidate not in seen:
|
| 78 |
+
seen.add(candidate)
|
| 79 |
+
corpus.append(candidate)
|
| 80 |
+
return corpus
|
| 81 |
+
|
| 82 |
+
|
| 83 |
+
def evaluate_model(
|
| 84 |
+
model: Any,
|
| 85 |
+
rows: list[dict[str, Any]],
|
| 86 |
+
*,
|
| 87 |
+
batch_size: int = 32,
|
| 88 |
+
) -> dict[str, Any]:
|
| 89 |
+
corpus = build_global_corpus(rows)
|
| 90 |
+
document_indexes = {document: index for index, document in enumerate(corpus)}
|
| 91 |
+
query_texts = [f"query: {row['query']}" for row in rows]
|
| 92 |
+
passage_texts = [f"passage: {document}" for document in corpus]
|
| 93 |
+
|
| 94 |
+
query_embeddings = model.encode(
|
| 95 |
+
query_texts,
|
| 96 |
+
batch_size=batch_size,
|
| 97 |
+
normalize_embeddings=True,
|
| 98 |
+
convert_to_numpy=True,
|
| 99 |
+
show_progress_bar=True,
|
| 100 |
+
)
|
| 101 |
+
passage_embeddings = model.encode(
|
| 102 |
+
passage_texts,
|
| 103 |
+
batch_size=batch_size,
|
| 104 |
+
normalize_embeddings=True,
|
| 105 |
+
convert_to_numpy=True,
|
| 106 |
+
show_progress_bar=True,
|
| 107 |
+
)
|
| 108 |
+
scores = np.asarray(query_embeddings) @ np.asarray(passage_embeddings).T
|
| 109 |
+
positive_indexes = [document_indexes[row["positive"]] for row in rows]
|
| 110 |
+
ranks = calculate_ranks(scores, positive_indexes)
|
| 111 |
+
tags = [row["challenge_tags"] for row in rows]
|
| 112 |
+
return summarize_ranks(ranks, tags=tags, documents=len(corpus))
|
| 113 |
+
|
| 114 |
+
|
| 115 |
+
def calculate_ranks(
|
| 116 |
+
scores: np.ndarray,
|
| 117 |
+
positive_indexes: list[int],
|
| 118 |
+
) -> list[int]:
|
| 119 |
+
if scores.ndim != 2:
|
| 120 |
+
raise ValueError("scores must be a two-dimensional query-document matrix")
|
| 121 |
+
if scores.shape[0] != len(positive_indexes):
|
| 122 |
+
raise ValueError("scores and positive_indexes must contain the same queries")
|
| 123 |
+
|
| 124 |
+
ranks: list[int] = []
|
| 125 |
+
for query_index, positive_index in enumerate(positive_indexes):
|
| 126 |
+
if positive_index < 0 or positive_index >= scores.shape[1]:
|
| 127 |
+
raise ValueError(f"Invalid positive index for query {query_index}")
|
| 128 |
+
order = np.argsort(-scores[query_index], kind="stable")
|
| 129 |
+
position = np.flatnonzero(order == positive_index)
|
| 130 |
+
ranks.append(int(position[0]) + 1)
|
| 131 |
+
return ranks
|
| 132 |
+
|
| 133 |
+
|
| 134 |
+
def summarize_ranks(
|
| 135 |
+
ranks: list[int],
|
| 136 |
+
*,
|
| 137 |
+
tags: list[list[str]],
|
| 138 |
+
documents: int,
|
| 139 |
+
) -> dict[str, Any]:
|
| 140 |
+
if not ranks or len(ranks) != len(tags):
|
| 141 |
+
raise ValueError("ranks and tags must be non-empty and have equal length")
|
| 142 |
+
|
| 143 |
+
overall = _rank_metrics(ranks)
|
| 144 |
+
tag_ranks: defaultdict[str, list[int]] = defaultdict(list)
|
| 145 |
+
for rank, row_tags in zip(ranks, tags):
|
| 146 |
+
for tag in set(row_tags):
|
| 147 |
+
tag_ranks[tag].append(rank)
|
| 148 |
+
|
| 149 |
+
return {
|
| 150 |
+
"queries": len(ranks),
|
| 151 |
+
"documents": documents,
|
| 152 |
+
**overall,
|
| 153 |
+
"per_challenge": {
|
| 154 |
+
tag: {"queries": len(values), **_rank_metrics(values)}
|
| 155 |
+
for tag, values in sorted(tag_ranks.items())
|
| 156 |
+
},
|
| 157 |
+
"ranks": ranks,
|
| 158 |
+
}
|
| 159 |
+
|
| 160 |
+
|
| 161 |
+
def _rank_metrics(ranks: list[int]) -> dict[str, float]:
|
| 162 |
+
values = np.asarray(ranks, dtype=np.float64)
|
| 163 |
+
discounted_gains_at_10 = np.where(
|
| 164 |
+
values <= 10,
|
| 165 |
+
1.0 / np.log2(values + 1.0),
|
| 166 |
+
0.0,
|
| 167 |
+
)
|
| 168 |
+
return {
|
| 169 |
+
"top1": float(np.mean(values <= 1)),
|
| 170 |
+
"recall_at_3": float(np.mean(values <= 3)),
|
| 171 |
+
"recall_at_5": float(np.mean(values <= 5)),
|
| 172 |
+
"recall_at_10": float(np.mean(values <= 10)),
|
| 173 |
+
"mrr": float(np.mean(1.0 / values)),
|
| 174 |
+
"ndcg_at_10": float(np.mean(discounted_gains_at_10)),
|
| 175 |
+
"mean_rank": float(np.mean(values)),
|
| 176 |
+
}
|
| 177 |
+
|
| 178 |
+
|
| 179 |
+
def _required_text(value: Any, line_number: int, field: str) -> str:
|
| 180 |
+
if not isinstance(value, str) or not value.strip():
|
| 181 |
+
raise ValueError(f"Line {line_number}: {field} must be a non-empty string")
|
| 182 |
+
return " ".join(value.split())
|
| 183 |
+
|
| 184 |
+
|
| 185 |
+
def _parse_model_specs(values: list[str]) -> list[tuple[str, str]]:
|
| 186 |
+
if not values:
|
| 187 |
+
return [("base", DEFAULT_BASE_MODEL)]
|
| 188 |
+
parsed: list[tuple[str, str]] = []
|
| 189 |
+
labels: set[str] = set()
|
| 190 |
+
for value in values:
|
| 191 |
+
if "=" not in value:
|
| 192 |
+
raise ValueError("Each --model must use LABEL=MODEL_OR_PATH")
|
| 193 |
+
label, model_path = (part.strip() for part in value.split("=", 1))
|
| 194 |
+
if not label or not model_path:
|
| 195 |
+
raise ValueError("Each --model must use LABEL=MODEL_OR_PATH")
|
| 196 |
+
if label in labels:
|
| 197 |
+
raise ValueError(f"Duplicate model label: {label}")
|
| 198 |
+
labels.add(label)
|
| 199 |
+
parsed.append((label, model_path))
|
| 200 |
+
return parsed
|
| 201 |
+
|
| 202 |
+
|
| 203 |
+
def _parse_args() -> argparse.Namespace:
|
| 204 |
+
parser = argparse.ArgumentParser(description=__doc__)
|
| 205 |
+
parser.add_argument("--test-file", required=True)
|
| 206 |
+
parser.add_argument(
|
| 207 |
+
"--model",
|
| 208 |
+
action="append",
|
| 209 |
+
default=[],
|
| 210 |
+
metavar="LABEL=MODEL_OR_PATH",
|
| 211 |
+
help="Repeat to compare multiple models on exactly the same corpus.",
|
| 212 |
+
)
|
| 213 |
+
parser.add_argument("--device", default=None)
|
| 214 |
+
parser.add_argument("--batch-size", type=int, default=32)
|
| 215 |
+
parser.add_argument("--output-json")
|
| 216 |
+
return parser.parse_args()
|
| 217 |
+
|
| 218 |
+
|
| 219 |
+
def main() -> None:
|
| 220 |
+
args = _parse_args()
|
| 221 |
+
if args.batch_size < 1:
|
| 222 |
+
raise ValueError("--batch-size must be greater than zero")
|
| 223 |
+
try:
|
| 224 |
+
from sentence_transformers import SentenceTransformer
|
| 225 |
+
except ImportError as exc:
|
| 226 |
+
raise RuntimeError(
|
| 227 |
+
"Install requirements-training.txt before evaluating retrievers"
|
| 228 |
+
) from exc
|
| 229 |
+
|
| 230 |
+
rows = read_evaluation_rows(Path(args.test_file))
|
| 231 |
+
results: dict[str, Any] = {}
|
| 232 |
+
for label, model_path in _parse_model_specs(args.model):
|
| 233 |
+
print(f"Evaluating {label}: {model_path}")
|
| 234 |
+
model = SentenceTransformer(model_path, device=args.device)
|
| 235 |
+
results[label] = evaluate_model(
|
| 236 |
+
model,
|
| 237 |
+
rows,
|
| 238 |
+
batch_size=args.batch_size,
|
| 239 |
+
)
|
| 240 |
+
|
| 241 |
+
output = json.dumps(results, ensure_ascii=False, indent=2)
|
| 242 |
+
print(output)
|
| 243 |
+
if args.output_json:
|
| 244 |
+
output_path = Path(args.output_json)
|
| 245 |
+
output_path.parent.mkdir(parents=True, exist_ok=True)
|
| 246 |
+
output_path.write_text(output + "\n", encoding="utf-8")
|
| 247 |
+
|
| 248 |
+
|
| 249 |
+
if __name__ == "__main__":
|
| 250 |
+
main()
|
tests/test_evaluate_place_retriever.py
ADDED
|
@@ -0,0 +1,51 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from pathlib import Path
|
| 2 |
+
|
| 3 |
+
import numpy as np
|
| 4 |
+
|
| 5 |
+
from scripts.evaluate_place_retriever import (
|
| 6 |
+
build_global_corpus,
|
| 7 |
+
calculate_ranks,
|
| 8 |
+
read_evaluation_rows,
|
| 9 |
+
summarize_ranks,
|
| 10 |
+
)
|
| 11 |
+
|
| 12 |
+
|
| 13 |
+
ROOT = Path(__file__).resolve().parents[1]
|
| 14 |
+
TEST_DATA = ROOT / "data" / "training" / "places_retrieval_v1" / "test.jsonl"
|
| 15 |
+
VALIDATION_DATA = (
|
| 16 |
+
ROOT / "data" / "training" / "places_retrieval_v1" / "validation.jsonl"
|
| 17 |
+
)
|
| 18 |
+
|
| 19 |
+
|
| 20 |
+
def test_challenge_evaluation_uses_one_global_candidate_pool() -> None:
|
| 21 |
+
for path, expected_queries in ((VALIDATION_DATA, 60), (TEST_DATA, 80)):
|
| 22 |
+
rows = read_evaluation_rows(path)
|
| 23 |
+
corpus = build_global_corpus(rows)
|
| 24 |
+
|
| 25 |
+
assert len(rows) == expected_queries
|
| 26 |
+
assert len(corpus) == 20
|
| 27 |
+
assert set(corpus) == {row["positive"] for row in rows}
|
| 28 |
+
|
| 29 |
+
|
| 30 |
+
def test_rank_metrics_are_calculated_over_all_documents_and_by_challenge() -> None:
|
| 31 |
+
scores = np.asarray(
|
| 32 |
+
[
|
| 33 |
+
[0.9, 0.8, 0.1, 0.0],
|
| 34 |
+
[0.8, 0.9, 0.7, 0.1],
|
| 35 |
+
[0.4, 0.3, 0.2, 0.1],
|
| 36 |
+
]
|
| 37 |
+
)
|
| 38 |
+
ranks = calculate_ranks(scores, [0, 2, 3])
|
| 39 |
+
summary = summarize_ranks(
|
| 40 |
+
ranks,
|
| 41 |
+
tags=[["canonical"], ["colloquial"], ["colloquial", "misspelling"]],
|
| 42 |
+
documents=4,
|
| 43 |
+
)
|
| 44 |
+
|
| 45 |
+
assert ranks == [1, 3, 4]
|
| 46 |
+
assert summary["queries"] == 3
|
| 47 |
+
assert summary["documents"] == 4
|
| 48 |
+
assert summary["top1"] == 1 / 3
|
| 49 |
+
assert summary["recall_at_3"] == 2 / 3
|
| 50 |
+
assert summary["per_challenge"]["colloquial"]["queries"] == 2
|
| 51 |
+
assert summary["per_challenge"]["misspelling"]["mean_rank"] == 4.0
|
tests/test_place_training_datasets.py
CHANGED
|
@@ -4,6 +4,8 @@ from pathlib import Path
|
|
| 4 |
|
| 5 |
from scripts.build_place_training_datasets import (
|
| 6 |
INTENT_CONCEPTS,
|
|
|
|
|
|
|
| 7 |
_concept_key,
|
| 8 |
_expanded_place_specs,
|
| 9 |
_place_document,
|
|
@@ -52,7 +54,7 @@ def test_intent_dataset_is_valid_open_vocabulary_and_leak_free() -> None:
|
|
| 52 |
|
| 53 |
def test_retrieval_dataset_matches_training_contract_and_has_unique_positives() -> None:
|
| 54 |
seen_queries: set[str] = set()
|
| 55 |
-
|
| 56 |
counts: dict[str, int] = {}
|
| 57 |
|
| 58 |
for split in ("train", "validation", "test"):
|
|
@@ -62,21 +64,24 @@ def test_retrieval_dataset_matches_training_contract_and_has_unique_positives()
|
|
| 62 |
counts[split] = len(rows)
|
| 63 |
for row in rows:
|
| 64 |
assert row["query"].casefold() not in seen_queries
|
| 65 |
-
|
|
|
|
| 66 |
seen_queries.add(row["query"].casefold())
|
| 67 |
-
|
| 68 |
assert row["positive"].startswith("Nombre: ")
|
| 69 |
assert "Tipo registrado:" in row["positive"]
|
| 70 |
assert "Descripcion:" in row["positive"]
|
| 71 |
assert "Etiquetas:" in row["positive"]
|
| 72 |
-
|
| 73 |
-
assert len(
|
|
|
|
| 74 |
assert row["positive"] not in row["hard_negatives"]
|
| 75 |
assert not row["query"].startswith("query: ")
|
| 76 |
assert not row["positive"].startswith("passage: ")
|
| 77 |
|
| 78 |
-
assert counts == {"train": 160, "validation":
|
| 79 |
-
assert 150 <=
|
|
|
|
| 80 |
|
| 81 |
|
| 82 |
def test_retrieval_hard_negatives_never_use_a_sibling_of_the_positive() -> None:
|
|
@@ -88,20 +93,61 @@ def test_retrieval_hard_negatives_never_use_a_sibling_of_the_positive() -> None:
|
|
| 88 |
document_concepts = {
|
| 89 |
_place_document(spec): _concept_key(spec) for spec in specs
|
| 90 |
}
|
| 91 |
-
|
|
|
|
| 92 |
if split != "train":
|
| 93 |
assert len({_concept_key(spec) for spec in specs}) == len(specs)
|
| 94 |
-
|
| 95 |
-
|
| 96 |
-
|
| 97 |
assert all(
|
| 98 |
document_concepts[negative] != positive_concept
|
| 99 |
for negative in row["hard_negatives"]
|
| 100 |
)
|
| 101 |
|
| 102 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 103 |
def test_dataset_manifest_matches_generated_files() -> None:
|
| 104 |
manifest = json.loads((DATA / "manifest.json").read_text(encoding="utf-8"))
|
| 105 |
assert manifest["schema_version"] == 1
|
| 106 |
assert manifest["intent"]["train"]["records"] == 132
|
| 107 |
assert manifest["retrieval"]["train"]["records"] == 160
|
|
|
|
|
|
|
|
|
| 4 |
|
| 5 |
from scripts.build_place_training_datasets import (
|
| 6 |
INTENT_CONCEPTS,
|
| 7 |
+
RETRIEVAL_TEST_QUERY_VARIANTS,
|
| 8 |
+
RETRIEVAL_VALIDATION_QUERY_VARIANTS,
|
| 9 |
_concept_key,
|
| 10 |
_expanded_place_specs,
|
| 11 |
_place_document,
|
|
|
|
| 54 |
|
| 55 |
def test_retrieval_dataset_matches_training_contract_and_has_unique_positives() -> None:
|
| 56 |
seen_queries: set[str] = set()
|
| 57 |
+
positive_splits: dict[str, str] = {}
|
| 58 |
counts: dict[str, int] = {}
|
| 59 |
|
| 60 |
for split in ("train", "validation", "test"):
|
|
|
|
| 64 |
counts[split] = len(rows)
|
| 65 |
for row in rows:
|
| 66 |
assert row["query"].casefold() not in seen_queries
|
| 67 |
+
positive_owner = positive_splits.get(row["positive"])
|
| 68 |
+
assert positive_owner in (None, split)
|
| 69 |
seen_queries.add(row["query"].casefold())
|
| 70 |
+
positive_splits[row["positive"]] = split
|
| 71 |
assert row["positive"].startswith("Nombre: ")
|
| 72 |
assert "Tipo registrado:" in row["positive"]
|
| 73 |
assert "Descripcion:" in row["positive"]
|
| 74 |
assert "Etiquetas:" in row["positive"]
|
| 75 |
+
expected_negatives = 3 if split == "train" else 7
|
| 76 |
+
assert len(row["hard_negatives"]) == expected_negatives
|
| 77 |
+
assert len(set(row["hard_negatives"])) == expected_negatives
|
| 78 |
assert row["positive"] not in row["hard_negatives"]
|
| 79 |
assert not row["query"].startswith("query: ")
|
| 80 |
assert not row["positive"].startswith("passage: ")
|
| 81 |
|
| 82 |
+
assert counts == {"train": 160, "validation": 60, "test": 80}
|
| 83 |
+
assert 150 <= counts["train"] <= 200
|
| 84 |
+
assert len(positive_splits) == 200
|
| 85 |
|
| 86 |
|
| 87 |
def test_retrieval_hard_negatives_never_use_a_sibling_of_the_positive() -> None:
|
|
|
|
| 93 |
document_concepts = {
|
| 94 |
_place_document(spec): _concept_key(spec) for spec in specs
|
| 95 |
}
|
| 96 |
+
if split == "train":
|
| 97 |
+
assert len(rows) == len(specs)
|
| 98 |
if split != "train":
|
| 99 |
assert len({_concept_key(spec) for spec in specs}) == len(specs)
|
| 100 |
+
assert {row["positive"] for row in rows} == set(document_concepts)
|
| 101 |
+
for row in rows:
|
| 102 |
+
positive_concept = document_concepts[row["positive"]]
|
| 103 |
assert all(
|
| 104 |
document_concepts[negative] != positive_concept
|
| 105 |
for negative in row["hard_negatives"]
|
| 106 |
)
|
| 107 |
|
| 108 |
|
| 109 |
+
def test_retrieval_evaluation_has_curated_challenges_and_repeated_qrels() -> None:
|
| 110 |
+
assert len(RETRIEVAL_VALIDATION_QUERY_VARIANTS) == 20
|
| 111 |
+
assert len(RETRIEVAL_TEST_QUERY_VARIANTS) == 20
|
| 112 |
+
for split, expected_rows, expected_repetitions in (
|
| 113 |
+
("validation", 60, 3),
|
| 114 |
+
("test", 80, 4),
|
| 115 |
+
):
|
| 116 |
+
path = DATA / "places_retrieval_v1" / f"{split}.jsonl"
|
| 117 |
+
rows = [
|
| 118 |
+
json.loads(line)
|
| 119 |
+
for line in path.read_text(encoding="utf-8").splitlines()
|
| 120 |
+
]
|
| 121 |
+
positive_counts = Counter(row["positive"] for row in rows)
|
| 122 |
+
tag_counts = Counter(
|
| 123 |
+
tag
|
| 124 |
+
for row in rows
|
| 125 |
+
for tag in row["challenge_tags"]
|
| 126 |
+
)
|
| 127 |
+
|
| 128 |
+
assert len(rows) == expected_rows
|
| 129 |
+
assert set(positive_counts.values()) == {expected_repetitions}
|
| 130 |
+
assert tag_counts["canonical"] == 20
|
| 131 |
+
assert tag_counts["colloquial"] >= 10
|
| 132 |
+
assert tag_counts["implicit"] >= 20
|
| 133 |
+
|
| 134 |
+
test_rows = [
|
| 135 |
+
json.loads(line)
|
| 136 |
+
for line in (DATA / "places_retrieval_v1" / "test.jsonl")
|
| 137 |
+
.read_text(encoding="utf-8")
|
| 138 |
+
.splitlines()
|
| 139 |
+
]
|
| 140 |
+
test_tag_counts = Counter(
|
| 141 |
+
tag for row in test_rows for tag in row["challenge_tags"]
|
| 142 |
+
)
|
| 143 |
+
assert test_tag_counts["misspelling"] >= 10
|
| 144 |
+
assert test_tag_counts["contrastive"] >= 4
|
| 145 |
+
|
| 146 |
+
|
| 147 |
def test_dataset_manifest_matches_generated_files() -> None:
|
| 148 |
manifest = json.loads((DATA / "manifest.json").read_text(encoding="utf-8"))
|
| 149 |
assert manifest["schema_version"] == 1
|
| 150 |
assert manifest["intent"]["train"]["records"] == 132
|
| 151 |
assert manifest["retrieval"]["train"]["records"] == 160
|
| 152 |
+
assert manifest["retrieval"]["validation"]["records"] == 60
|
| 153 |
+
assert manifest["retrieval"]["test"]["records"] == 80
|