Standardisation 8
This commit is contained in:
+19
@@ -174,6 +174,8 @@ scripts migrés vers cette convention sont actuellement :
|
|||||||
- =page_splitter.py=, =cutleft.py=, =plotting.py= et
|
- =page_splitter.py=, =cutleft.py=, =plotting.py= et
|
||||||
=splitting_int.py= ;
|
=splitting_int.py= ;
|
||||||
- =gemini_for_labels.py= ;
|
- =gemini_for_labels.py= ;
|
||||||
|
- =correction.py=, =submit_batches.py=, =batch_status.py= et
|
||||||
|
=fetch_batched_results.py= ;
|
||||||
- =annotating.py=, =annotating_with_checks.py= et
|
- =annotating.py=, =annotating_with_checks.py= et
|
||||||
=annotating_by_label.py= ;
|
=annotating_by_label.py= ;
|
||||||
- =reading_annotations.py= et =reading_grouped_annotations.py= ;
|
- =reading_annotations.py= et =reading_grouped_annotations.py= ;
|
||||||
@@ -343,6 +345,13 @@ Optional : Set proxy with ~export HTTPS_PROXY="http://10.0.0.1:3128"~
|
|||||||
|
|
||||||
Fais les requêtes de correction à Gemini.
|
Fais les requêtes de correction à Gemini.
|
||||||
|
|
||||||
|
=correction.py= peut être relancé sans supprimer son état. Les
|
||||||
|
fichiers =correction.json= et =correction_progress.json= sont mis à
|
||||||
|
jour atomiquement. Avec =--overwrite=, leur version précédente reste
|
||||||
|
en place jusqu'à la première écriture réussie de la nouvelle
|
||||||
|
exécution. =--reset= est la seule option qui supprime explicitement
|
||||||
|
cet état et restaure les fichiers =*_old.pdf=.
|
||||||
|
|
||||||
L'argument =limit= limite le nombre de requêtes à Gemini Pro
|
L'argument =limit= limite le nombre de requêtes à Gemini Pro
|
||||||
(chères), pour une version low cost, passer =--limit 0=, toutes
|
(chères), pour une version low cost, passer =--limit 0=, toutes
|
||||||
les requêtes seront sur Gemini Flash.
|
les requêtes seront sur Gemini Flash.
|
||||||
@@ -357,6 +366,16 @@ Optional : Set proxy with ~export HTTPS_PROXY="http://10.0.0.1:3128"~
|
|||||||
+ =python batch_status.py=
|
+ =python batch_status.py=
|
||||||
+ =python fetch_batched_results.py Interro=
|
+ =python fetch_batched_results.py Interro=
|
||||||
+ =python correction.py Interro --deal-with-batched=
|
+ =python correction.py Interro --deal-with-batched=
|
||||||
|
|
||||||
|
Les quatre commandes de ce flux suivent la convention des scripts
|
||||||
|
standardisés. Les fichiers de requêtes et le résultat JSONL combiné
|
||||||
|
sont publiés atomiquement : une interruption ne laisse pas de fichier
|
||||||
|
final partiellement écrit. =submit_batches.py= conserve aussi les
|
||||||
|
identifiants distants dans =batch_jobs.json= ; la récupération les
|
||||||
|
utilise en priorité et garde la recherche par nom pour les anciens
|
||||||
|
batchs. =batch_status.py --download JOB --output
|
||||||
|
resultat.jsonl= permet aussi de télécharger atomiquement le résultat
|
||||||
|
d'un job particulier.
|
||||||
3. =python post-correction.py Interro=
|
3. =python post-correction.py Interro=
|
||||||
|
|
||||||
- Essaye de corriger des erreurs d'encodage/d'accents dans
|
- Essaye de corriger des erreurs d'encodage/d'accents dans
|
||||||
|
|||||||
+79
-74
@@ -1,88 +1,93 @@
|
|||||||
import os
|
from __future__ import annotations
|
||||||
import sys
|
|
||||||
import argparse
|
import argparse
|
||||||
|
from collections.abc import Sequence
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
from google import genai
|
from google import genai
|
||||||
|
|
||||||
if "GEMINI_API_KEY" not in os.environ:
|
import config
|
||||||
sys.exit("Error: GEMINI_API_KEY environment variable not set.")
|
from copienator import (
|
||||||
|
CliError,
|
||||||
client = genai.Client()
|
ExitCode,
|
||||||
|
atomic_write_bytes,
|
||||||
def list_jobs():
|
execute,
|
||||||
print("Fetching recent batch jobs...\n")
|
standard_parser,
|
||||||
try:
|
)
|
||||||
batch_jobs = client.batches.list()
|
|
||||||
jobs_found = False
|
|
||||||
|
|
||||||
for job in batch_jobs:
|
|
||||||
jobs_found = True
|
|
||||||
state = job.state.name if hasattr(job.state, 'name') else job.state
|
|
||||||
|
|
||||||
print("-" * 60)
|
|
||||||
print(f"Job Name: {job.name}")
|
|
||||||
|
|
||||||
if hasattr(job, 'display_name') and job.display_name:
|
|
||||||
print(f"Display Name: {job.display_name}")
|
|
||||||
|
|
||||||
print(f"State: {state}")
|
|
||||||
|
|
||||||
if state == 'JOB_STATE_FAILED' and hasattr(job, 'error'):
|
|
||||||
print(f"Error: {job.error}")
|
|
||||||
|
|
||||||
if state == 'JOB_STATE_SUCCEEDED' and hasattr(job, 'dest') and job.dest:
|
|
||||||
if hasattr(job.dest, 'file_name') and job.dest.file_name:
|
|
||||||
print(f"Output File: {job.dest.file_name}")
|
|
||||||
|
|
||||||
if not jobs_found:
|
|
||||||
print("No batch jobs found.")
|
|
||||||
else:
|
|
||||||
print("-" * 60)
|
|
||||||
print("\nTo download a completed job, run:")
|
|
||||||
print("python batch_status.py --download batches/<YOUR_BATCH_ID>")
|
|
||||||
|
|
||||||
except Exception as e:
|
|
||||||
sys.exit(f"An error occurred while listing jobs: {e}")
|
|
||||||
|
|
||||||
|
|
||||||
def download_job(job_name):
|
def _client():
|
||||||
print(f"Checking status for {job_name}...\n")
|
if not config.API_KEY:
|
||||||
try:
|
raise CliError("GEMINI_API_KEY is not configured")
|
||||||
job = client.batches.get(name=job_name)
|
return genai.Client(api_key=config.API_KEY)
|
||||||
state = job.state.name if hasattr(job.state, 'name') else job.state
|
|
||||||
|
|
||||||
print(f"State: {state}")
|
|
||||||
|
|
||||||
if state != 'JOB_STATE_SUCCEEDED':
|
def list_jobs(*, client=None) -> ExitCode:
|
||||||
print("Job is not ready yet or has failed.")
|
client = client or _client()
|
||||||
if state == 'JOB_STATE_FAILED' and hasattr(job, 'error'):
|
print("Fetching recent batch jobs...")
|
||||||
print(f"Error: {job.error}")
|
jobs = list(client.batches.list())
|
||||||
return
|
for job in jobs:
|
||||||
|
state = job.state.name if hasattr(job.state, "name") else job.state
|
||||||
|
print(f"{job.name}: {state}")
|
||||||
|
if getattr(job, "display_name", None):
|
||||||
|
print(f" Display name: {job.display_name}")
|
||||||
|
if state == "JOB_STATE_FAILED" and getattr(job, "error", None):
|
||||||
|
print(f" Error: {job.error}")
|
||||||
|
destination = getattr(job, "dest", None)
|
||||||
|
if state == "JOB_STATE_SUCCEEDED" and getattr(
|
||||||
|
destination, "file_name", None
|
||||||
|
):
|
||||||
|
print(f" Output file: {destination.file_name}")
|
||||||
|
if not jobs:
|
||||||
|
print("No batch jobs found.")
|
||||||
|
return ExitCode.SUCCESS
|
||||||
|
|
||||||
if hasattr(job, 'dest') and job.dest and hasattr(job.dest, 'file_name') and job.dest.file_name:
|
|
||||||
result_file_name = job.dest.file_name
|
|
||||||
print(f"Downloading results from {result_file_name}...")
|
|
||||||
|
|
||||||
file_content_bytes = client.files.download(file=result_file_name)
|
def download_job(
|
||||||
output_path = f"results_{job_name.replace('/', '_')}.jsonl"
|
job_name: str,
|
||||||
|
*,
|
||||||
|
output: Path | None = None,
|
||||||
|
client=None,
|
||||||
|
) -> ExitCode:
|
||||||
|
client = client or _client()
|
||||||
|
job = client.batches.get(name=job_name)
|
||||||
|
state = job.state.name if hasattr(job.state, "name") else job.state
|
||||||
|
print(f"State: {state}")
|
||||||
|
if state != "JOB_STATE_SUCCEEDED":
|
||||||
|
if state == "JOB_STATE_FAILED" and getattr(job, "error", None):
|
||||||
|
print(f"Error: {job.error}")
|
||||||
|
return ExitCode.PARTIAL
|
||||||
|
destination = getattr(job, "dest", None)
|
||||||
|
file_name = getattr(destination, "file_name", None)
|
||||||
|
if not file_name:
|
||||||
|
print("Job succeeded but no output file was found.")
|
||||||
|
return ExitCode.PARTIAL
|
||||||
|
payload = client.files.download(file=file_name)
|
||||||
|
output_path = output or Path(f"results_{job_name.replace('/', '_')}.jsonl")
|
||||||
|
atomic_write_bytes(output_path, payload)
|
||||||
|
print(f"Saved batch results to {output_path}")
|
||||||
|
return ExitCode.SUCCESS
|
||||||
|
|
||||||
with open(output_path, "wb") as f:
|
|
||||||
f.write(file_content_bytes)
|
|
||||||
|
|
||||||
print(f"Success! Saved to {output_path}")
|
def build_parser() -> argparse.ArgumentParser:
|
||||||
print(f"You can now feed this to your correction script using: --deal-with-batched {output_path}")
|
parser = standard_parser("List or download Gemini correction batch jobs")
|
||||||
else:
|
parser.add_argument("--download", metavar="JOB_NAME")
|
||||||
print("Job succeeded but no output file was found.")
|
parser.add_argument("--output", type=Path, help="Downloaded JSONL destination")
|
||||||
|
return parser
|
||||||
|
|
||||||
|
|
||||||
|
def main(argv: Sequence[str] | None = None) -> int:
|
||||||
|
parser = build_parser()
|
||||||
|
|
||||||
|
def handle(args: argparse.Namespace) -> ExitCode:
|
||||||
|
if args.output is not None and not args.download:
|
||||||
|
raise CliError("--output requires --download", ExitCode.INVALID_ARGUMENTS)
|
||||||
|
if args.download:
|
||||||
|
return download_job(args.download, output=args.output)
|
||||||
|
return list_jobs()
|
||||||
|
|
||||||
|
return execute(parser, argv, handle)
|
||||||
|
|
||||||
except Exception as e:
|
|
||||||
sys.exit(f"An error occurred while fetching the job: {e}")
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
parser = argparse.ArgumentParser(description="Manage Gemini Batch Jobs")
|
raise SystemExit(main())
|
||||||
parser.add_argument("--download", type=str, metavar="JOB_NAME",
|
|
||||||
help="Download the results for a specific batch job (e.g. batches/123456)")
|
|
||||||
args = parser.parse_args()
|
|
||||||
|
|
||||||
if args.download:
|
|
||||||
download_job(args.download)
|
|
||||||
else:
|
|
||||||
list_jobs()
|
|
||||||
|
|||||||
@@ -14,6 +14,7 @@ from .cli import (
|
|||||||
from .json_io import (
|
from .json_io import (
|
||||||
JsonLockTimeout,
|
JsonLockTimeout,
|
||||||
atomic_update_json,
|
atomic_update_json,
|
||||||
|
atomic_write_bytes,
|
||||||
atomic_write_json,
|
atomic_write_json,
|
||||||
atomic_write_text,
|
atomic_write_text,
|
||||||
read_json,
|
read_json,
|
||||||
@@ -32,6 +33,7 @@ __all__ = [
|
|||||||
"WorkspaceNotFoundError",
|
"WorkspaceNotFoundError",
|
||||||
"WorkspaceValidationError",
|
"WorkspaceValidationError",
|
||||||
"atomic_update_json",
|
"atomic_update_json",
|
||||||
|
"atomic_write_bytes",
|
||||||
"atomic_write_json",
|
"atomic_write_json",
|
||||||
"atomic_write_text",
|
"atomic_write_text",
|
||||||
"evaluation_parser",
|
"evaluation_parser",
|
||||||
|
|||||||
@@ -78,6 +78,10 @@ def atomic_write_text(
|
|||||||
_atomic_write(Path(path), text.encode(encoding))
|
_atomic_write(Path(path), text.encode(encoding))
|
||||||
|
|
||||||
|
|
||||||
|
def atomic_write_bytes(path: str | Path, payload: bytes) -> None:
|
||||||
|
_atomic_write(Path(path), payload)
|
||||||
|
|
||||||
|
|
||||||
def atomic_write_json(
|
def atomic_write_json(
|
||||||
path: str | Path,
|
path: str | Path,
|
||||||
value: JsonValue,
|
value: JsonValue,
|
||||||
|
|||||||
@@ -101,6 +101,14 @@ class EvaluationWorkspace:
|
|||||||
def correction_progress_file(self) -> Path:
|
def correction_progress_file(self) -> Path:
|
||||||
return self.root / "correction_progress.json"
|
return self.root / "correction_progress.json"
|
||||||
|
|
||||||
|
@property
|
||||||
|
def batch_jobs_file(self) -> Path:
|
||||||
|
return self.root / "batch_jobs.json"
|
||||||
|
|
||||||
|
@property
|
||||||
|
def batched_correction_result_file(self) -> Path:
|
||||||
|
return self.root / "batched_correction_result.jsonl"
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def manual_resolutions_file(self) -> Path:
|
def manual_resolutions_file(self) -> Path:
|
||||||
return self.root / "manual_resolutions.txt"
|
return self.root / "manual_resolutions.txt"
|
||||||
|
|||||||
+336
-208
@@ -1,91 +1,59 @@
|
|||||||
import sys
|
from __future__ import annotations
|
||||||
import os
|
|
||||||
import time
|
|
||||||
from pathlib import Path
|
|
||||||
import argparse
|
import argparse
|
||||||
import prompting
|
|
||||||
import signal
|
|
||||||
from google import genai
|
|
||||||
import base64
|
import base64
|
||||||
import shlex
|
|
||||||
import json
|
|
||||||
import threading
|
|
||||||
import concurrent.futures
|
import concurrent.futures
|
||||||
|
import json
|
||||||
|
import os
|
||||||
|
import shlex
|
||||||
|
import shutil
|
||||||
|
import sys
|
||||||
|
import threading
|
||||||
|
import time
|
||||||
|
from collections.abc import Sequence
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
from copienator import atomic_write_json
|
from google import genai
|
||||||
|
|
||||||
if len(sys.argv) < 2:
|
|
||||||
sys.exit("Usage: python script.py 'InterroTest/Ex 2/Group_1.jpg' OR <InputDir> OR 'file1' 'file2'")
|
|
||||||
|
|
||||||
# Parse Arguments
|
|
||||||
parser = argparse.ArgumentParser()
|
|
||||||
parser.add_argument("paths", nargs="+", help="List of images or directories")
|
|
||||||
parser.add_argument("--overwrite", action="store_true",
|
|
||||||
help="Force redo requests even if output exists")
|
|
||||||
parser.add_argument("--limit", type=int, help="limit calls to gemini rpo integer")
|
|
||||||
parser.add_argument("--refaire", action="store_true",
|
|
||||||
help="Redo specific copies/labels defined in refaire.json")
|
|
||||||
parser.add_argument("--batch", action="store_true",
|
|
||||||
help="Generate a JSONL file of requests to send to the Gemini Batch API")
|
|
||||||
parser.add_argument("--batch-from", type=str, metavar="LABEL",
|
|
||||||
help="Do live requests before LABEL, and batch requests from LABEL onwards")
|
|
||||||
parser.add_argument("--deal-with-batched", action="store_true",
|
|
||||||
help="Process a JSONL file containing completed batch results")
|
|
||||||
parser.add_argument("--reset", action="store_true",
|
|
||||||
help="Remove correction.json, revert _old.pdf, delete _new.pdf, then exit")
|
|
||||||
args, _ = parser.parse_known_args()
|
|
||||||
|
|
||||||
tasks = [] # List of tuples: (filepath_str, label_str)
|
|
||||||
results = {}
|
|
||||||
|
|
||||||
|
|
||||||
for path_str in args.paths:
|
|
||||||
arg_path = Path(path_str)
|
|
||||||
|
|
||||||
if not arg_path.exists():
|
|
||||||
print(f"Warning: {path_str} not found. Skipping.")
|
|
||||||
continue
|
|
||||||
|
|
||||||
if arg_path.is_file() and arg_path.suffix.lower() == ".jpg":
|
|
||||||
# Handle individual file
|
|
||||||
# Note: assumes structure InterroTest/Ex 2/Group_1.jpg
|
|
||||||
label = arg_path.parent.name
|
|
||||||
INPUT_DIR = arg_path.parent.parent.parent
|
|
||||||
COPIES_DIR = INPUT_DIR / "Copies"
|
|
||||||
GROUPS_DIR = INPUT_DIR / "Par label"
|
|
||||||
tasks.append((str(arg_path), label))
|
|
||||||
if label not in results:
|
|
||||||
results[label] = []
|
|
||||||
|
|
||||||
elif arg_path.is_dir():
|
|
||||||
INPUT_DIR = arg_path
|
|
||||||
COPIES_DIR = INPUT_DIR / "Copies"
|
|
||||||
GROUPS_DIR = INPUT_DIR / "Par label"
|
|
||||||
# Handle directory (original behavior)
|
|
||||||
for sub in GROUPS_DIR.iterdir():
|
|
||||||
if sub.is_dir():
|
|
||||||
label = sub.name
|
|
||||||
if label not in results:
|
|
||||||
results[label] = []
|
|
||||||
for img in sub.glob("*.jpg"):
|
|
||||||
tasks.append((str(img), label))
|
|
||||||
|
|
||||||
|
import config
|
||||||
|
import grouping
|
||||||
|
import prompting
|
||||||
|
from copienator import (
|
||||||
|
CliError,
|
||||||
|
EvaluationWorkspace,
|
||||||
|
ExitCode,
|
||||||
|
atomic_write_json,
|
||||||
|
atomic_write_text,
|
||||||
|
execute,
|
||||||
|
read_json,
|
||||||
|
target_parser,
|
||||||
|
workspace_from_target,
|
||||||
|
)
|
||||||
|
from utils import enonce_total, read_all_labels
|
||||||
|
|
||||||
NB_THREADS = 12
|
NB_THREADS = 12
|
||||||
|
|
||||||
# PROXY_URL = "http://192.168.241.1:3128"
|
# PROXY_URL = "http://192.168.241.1:3128"
|
||||||
PROXY_URL = None
|
|
||||||
|
|
||||||
if PROXY_URL:
|
|
||||||
os.environ["http_proxy"] = PROXY_URL
|
|
||||||
os.environ["https_proxy"] = PROXY_URL
|
|
||||||
|
|
||||||
import config
|
|
||||||
|
|
||||||
MODEL_ID_pro = config.MODEL_PRO_ID
|
MODEL_ID_pro = config.MODEL_PRO_ID
|
||||||
MODEL_ID_flash = config.MODEL_FLASH_ID
|
MODEL_ID_flash = config.MODEL_FLASH_ID
|
||||||
api_key = config.API_KEY
|
api_key = config.API_KEY
|
||||||
|
|
||||||
|
# Runtime globals retained while the processing helpers are migrated incrementally.
|
||||||
|
INPUT_DIR = Path()
|
||||||
|
COPIES_DIR = Path()
|
||||||
|
GROUPS_DIR = Path()
|
||||||
|
output_path = Path()
|
||||||
|
progress_path = Path()
|
||||||
|
tasks: list[tuple] = []
|
||||||
|
tasks_to_process: list[tuple] = []
|
||||||
|
results: dict = {}
|
||||||
|
completed_tasks: list = []
|
||||||
|
errors_summary: list = []
|
||||||
|
overwrite = False
|
||||||
|
limit = None
|
||||||
|
client = None
|
||||||
|
start_time = 0.0
|
||||||
|
|
||||||
# --- Thread-safe Logging ---
|
# --- Thread-safe Logging ---
|
||||||
log_lock = threading.Lock()
|
log_lock = threading.Lock()
|
||||||
thread_logs = {}
|
thread_logs = {}
|
||||||
@@ -113,57 +81,6 @@ def flush_thread_log(tid=None):
|
|||||||
f.write("\n".join(thread_logs[tid]) + "\n\n")
|
f.write("\n".join(thread_logs[tid]) + "\n\n")
|
||||||
thread_logs[tid].clear()
|
thread_logs[tid].clear()
|
||||||
|
|
||||||
def handle_interrupt(sig, frame):
|
|
||||||
"""Flush all partial/unfinished logs if program is interrupted."""
|
|
||||||
print("\nInterrupt received. Flushing partial logs...", file=sys.stderr)
|
|
||||||
for tid in list(thread_logs.keys()):
|
|
||||||
flush_thread_log(tid)
|
|
||||||
sys.exit(1)
|
|
||||||
|
|
||||||
|
|
||||||
signal.signal(signal.SIGINT, handle_interrupt)
|
|
||||||
signal.signal(signal.SIGTERM, handle_interrupt)
|
|
||||||
# ---------------------------
|
|
||||||
|
|
||||||
client = genai.Client(api_key=api_key)
|
|
||||||
output_path = INPUT_DIR / "correction.json"
|
|
||||||
progress_path = INPUT_DIR / "correction_progress.json"
|
|
||||||
|
|
||||||
if args.reset:
|
|
||||||
print("--- Running Reset ---")
|
|
||||||
if output_path.exists():
|
|
||||||
output_path.unlink()
|
|
||||||
print(f"Deleted: {output_path}")
|
|
||||||
if progress_path.exists():
|
|
||||||
progress_path.unlink()
|
|
||||||
print(f"Deleted: {progress_path}")
|
|
||||||
|
|
||||||
if COPIES_DIR.exists():
|
|
||||||
for copie_dir in COPIES_DIR.iterdir():
|
|
||||||
if not copie_dir.is_dir():
|
|
||||||
continue
|
|
||||||
|
|
||||||
# Revert _old.pdf files
|
|
||||||
for old_pdf in copie_dir.glob("*_old.pdf"):
|
|
||||||
orig_pdf = copie_dir / old_pdf.name.replace("_old.pdf", ".pdf")
|
|
||||||
if orig_pdf.exists():
|
|
||||||
orig_pdf.unlink() # Prevent FileExistsError on Windows
|
|
||||||
old_pdf.rename(orig_pdf)
|
|
||||||
print(f"Moved: {copie_dir.name}/{old_pdf.name} -> {orig_pdf.name}")
|
|
||||||
|
|
||||||
# Delete _new.pdf files
|
|
||||||
for new_pdf in copie_dir.glob("*_new.pdf"):
|
|
||||||
new_pdf.unlink()
|
|
||||||
print(f"Deleted: {copie_dir.name}/{new_pdf.name}")
|
|
||||||
|
|
||||||
sys.exit("Reset almost complete. For each deleted `_new`, you should manually delete the group in `Par label`")
|
|
||||||
|
|
||||||
start_time = time.time()
|
|
||||||
overwrite = args.overwrite
|
|
||||||
limit = args.limit
|
|
||||||
completed_tasks = []
|
|
||||||
errors_summary = []
|
|
||||||
|
|
||||||
# --- Lock for thread-safe file writing ---
|
# --- Lock for thread-safe file writing ---
|
||||||
io_lock = threading.Lock()
|
io_lock = threading.Lock()
|
||||||
pro_lock = threading.Lock()
|
pro_lock = threading.Lock()
|
||||||
@@ -171,21 +88,119 @@ pro_count = 0
|
|||||||
flash_count = 0
|
flash_count = 0
|
||||||
pro_quota_exhausted = False
|
pro_quota_exhausted = False
|
||||||
|
|
||||||
if overwrite:
|
|
||||||
if output_path.exists():
|
|
||||||
output_path.unlink()
|
|
||||||
if progress_path.exists():
|
|
||||||
progress_path.unlink()
|
|
||||||
else:
|
|
||||||
if progress_path.exists():
|
|
||||||
with open(progress_path, "r", encoding="utf-8") as f:
|
|
||||||
completed_tasks = json.load(f)
|
|
||||||
if output_path.exists():
|
|
||||||
with open(output_path, "r", encoding="utf-8") as f:
|
|
||||||
results = json.load(f)
|
|
||||||
|
|
||||||
completed_set = set((str(f), l) for f, l in completed_tasks)
|
def discover_tasks(
|
||||||
tasks_to_process = [t for t in tasks if (str(t[0]), t[1]) not in completed_set]
|
workspace: EvaluationWorkspace,
|
||||||
|
targets: list[Path],
|
||||||
|
) -> tuple[list[tuple[str, str]], list[str]]:
|
||||||
|
workspace.require_directories("Copies", "Par label")
|
||||||
|
discovered: list[tuple[str, str]] = []
|
||||||
|
warnings: list[str] = []
|
||||||
|
for target in targets:
|
||||||
|
if target.is_file():
|
||||||
|
if target.suffix.casefold() != ".jpg":
|
||||||
|
raise CliError(
|
||||||
|
f"Correction target is not a group JPG: {target}",
|
||||||
|
ExitCode.INVALID_ARGUMENTS,
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
target.relative_to(workspace.groups_dir)
|
||||||
|
except ValueError as exc:
|
||||||
|
raise CliError(
|
||||||
|
f"Group image is not inside {workspace.groups_dir}: {target}",
|
||||||
|
ExitCode.INVALID_ARGUMENTS,
|
||||||
|
) from exc
|
||||||
|
discovered.append((str(target), target.parent.name))
|
||||||
|
continue
|
||||||
|
|
||||||
|
group_directories = sorted(
|
||||||
|
(path for path in workspace.groups_dir.iterdir() if path.is_dir()),
|
||||||
|
key=lambda path: path.name.casefold(),
|
||||||
|
)
|
||||||
|
for group_directory in group_directories:
|
||||||
|
images = sorted(
|
||||||
|
group_directory.glob("*.jpg"), key=lambda path: path.name.casefold()
|
||||||
|
)
|
||||||
|
discovered.extend(
|
||||||
|
(str(image), group_directory.name) for image in images
|
||||||
|
)
|
||||||
|
if not group_directories:
|
||||||
|
warnings.append(f"No label groups found in {workspace.groups_dir}")
|
||||||
|
return list(dict.fromkeys(discovered)), warnings
|
||||||
|
|
||||||
|
|
||||||
|
def configure_runtime(
|
||||||
|
workspace: EvaluationWorkspace,
|
||||||
|
discovered_tasks: list[tuple[str, str]],
|
||||||
|
args: argparse.Namespace,
|
||||||
|
*,
|
||||||
|
api_client=None,
|
||||||
|
) -> None:
|
||||||
|
global INPUT_DIR, COPIES_DIR, GROUPS_DIR, output_path, progress_path
|
||||||
|
global tasks, tasks_to_process, results, completed_tasks, errors_summary
|
||||||
|
global overwrite, limit, client, start_time
|
||||||
|
global pro_count, flash_count, pro_quota_exhausted
|
||||||
|
|
||||||
|
INPUT_DIR = workspace.root
|
||||||
|
COPIES_DIR = workspace.copies_dir
|
||||||
|
GROUPS_DIR = workspace.groups_dir
|
||||||
|
output_path = workspace.correction_file
|
||||||
|
progress_path = workspace.correction_progress_file
|
||||||
|
tasks = list(discovered_tasks)
|
||||||
|
overwrite = bool(args.overwrite)
|
||||||
|
limit = args.limit
|
||||||
|
start_time = time.time()
|
||||||
|
errors_summary = []
|
||||||
|
completed_tasks = []
|
||||||
|
results = {label: [] for _file, label in tasks}
|
||||||
|
thread_logs.clear()
|
||||||
|
pro_count = 0
|
||||||
|
flash_count = 0
|
||||||
|
pro_quota_exhausted = False
|
||||||
|
|
||||||
|
if not overwrite:
|
||||||
|
if progress_path.is_file():
|
||||||
|
loaded_progress = read_json(progress_path)
|
||||||
|
if not isinstance(loaded_progress, list):
|
||||||
|
raise TypeError("correction_progress.json must contain a JSON array")
|
||||||
|
completed_tasks = loaded_progress
|
||||||
|
if output_path.is_file():
|
||||||
|
loaded_results = read_json(output_path)
|
||||||
|
if not isinstance(loaded_results, dict):
|
||||||
|
raise TypeError("correction.json must contain a JSON object")
|
||||||
|
results = loaded_results
|
||||||
|
|
||||||
|
completed_set = {(str(file_path), label) for file_path, label in completed_tasks}
|
||||||
|
tasks_to_process = [
|
||||||
|
task for task in tasks if (str(task[0]), task[1]) not in completed_set
|
||||||
|
]
|
||||||
|
client = api_client
|
||||||
|
|
||||||
|
|
||||||
|
def reset_workspace(workspace: EvaluationWorkspace) -> None:
|
||||||
|
"""Apply the explicitly requested correction reset."""
|
||||||
|
print("--- Running Reset ---")
|
||||||
|
for path in (workspace.correction_file, workspace.correction_progress_file):
|
||||||
|
if path.exists():
|
||||||
|
path.unlink()
|
||||||
|
print(f"Deleted: {path}")
|
||||||
|
if workspace.copies_dir.is_dir():
|
||||||
|
for copy_directory in workspace.copies_dir.iterdir():
|
||||||
|
if not copy_directory.is_dir():
|
||||||
|
continue
|
||||||
|
for old_pdf in copy_directory.glob("*_old.pdf"):
|
||||||
|
original = old_pdf.with_name(old_pdf.name.replace("_old.pdf", ".pdf"))
|
||||||
|
if original.exists():
|
||||||
|
original.unlink()
|
||||||
|
old_pdf.replace(original)
|
||||||
|
print(f"Moved: {copy_directory.name}/{old_pdf.name} -> {original.name}")
|
||||||
|
for new_pdf in copy_directory.glob("*_new.pdf"):
|
||||||
|
new_pdf.unlink()
|
||||||
|
print(f"Deleted: {copy_directory.name}/{new_pdf.name}")
|
||||||
|
print(
|
||||||
|
"Reset almost complete. Manually remove groups associated with deleted "
|
||||||
|
"_new PDFs from 'Par label'."
|
||||||
|
)
|
||||||
|
|
||||||
def call_gemini_with_retries(model_id, contents, config,
|
def call_gemini_with_retries(model_id, contents, config,
|
||||||
fallback_model_id=MODEL_ID_flash):
|
fallback_model_id=MODEL_ID_flash):
|
||||||
@@ -267,9 +282,6 @@ def correct_boxes_with_gemini(pid, label, pdf_path, original_feedbacks,
|
|||||||
|
|
||||||
return global_feedbacks + corrected_feedbacks
|
return global_feedbacks + corrected_feedbacks
|
||||||
|
|
||||||
import shutil
|
|
||||||
import grouping
|
|
||||||
|
|
||||||
def get_next_group_idx(label):
|
def get_next_group_idx(label):
|
||||||
"""Finds the next available Group index for a given label."""
|
"""Finds the next available Group index for a given label."""
|
||||||
target_folder = GROUPS_DIR / label
|
target_folder = GROUPS_DIR / label
|
||||||
@@ -278,8 +290,6 @@ def get_next_group_idx(label):
|
|||||||
if not existing: return 0
|
if not existing: return 0
|
||||||
return max([int(f.stem.split("_")[1]) for f in existing])
|
return max([int(f.stem.split("_")[1]) for f in existing])
|
||||||
|
|
||||||
from utils import read_all_labels, enonce_total
|
|
||||||
|
|
||||||
def handle_label_errors(pid, label, res, pdf_path):
|
def handle_label_errors(pid, label, res, pdf_path):
|
||||||
"""Handles Gemini labeling errors, moves/copies files, and returns new tasks."""
|
"""Handles Gemini labeling errors, moves/copies files, and returns new tasks."""
|
||||||
new_tasks = []
|
new_tasks = []
|
||||||
@@ -332,7 +342,7 @@ def handle_label_errors(pid, label, res, pdf_path):
|
|||||||
tprint(f"\tHandling additional-answer for {pid} {label}")
|
tprint(f"\tHandling additional-answer for {pid} {label}")
|
||||||
try:
|
try:
|
||||||
add_labels = json.loads(call_gemini_with_retries(MODEL_ID_flash, contents, config))
|
add_labels = json.loads(call_gemini_with_retries(MODEL_ID_flash, contents, config))
|
||||||
except Exception:
|
except Exception: # noqa: BLE001 - invalid auxiliary model response
|
||||||
add_labels = []
|
add_labels = []
|
||||||
|
|
||||||
keep_error = False
|
keep_error = False
|
||||||
@@ -374,7 +384,7 @@ def handle_label_errors(pid, label, res, pdf_path):
|
|||||||
|
|
||||||
def process_single_task(task_tuple, precomputed_response=None):
|
def process_single_task(task_tuple, precomputed_response=None):
|
||||||
try:
|
try:
|
||||||
global pro_count, flash_count, pro_quota_exhausted
|
global pro_count, flash_count
|
||||||
file_path = task_tuple[0]
|
file_path = task_tuple[0]
|
||||||
label = task_tuple[1]
|
label = task_tuple[1]
|
||||||
can_spawn_tasks = task_tuple[2] if len(task_tuple) > 2 else True
|
can_spawn_tasks = task_tuple[2] if len(task_tuple) > 2 else True
|
||||||
@@ -383,8 +393,7 @@ def process_single_task(task_tuple, precomputed_response=None):
|
|||||||
json_path = group_name + '.json'
|
json_path = group_name + '.json'
|
||||||
new_tasks = []
|
new_tasks = []
|
||||||
|
|
||||||
with open(json_path, 'r') as f:
|
group_data = read_json(json_path)
|
||||||
group_data = json.load(f)
|
|
||||||
|
|
||||||
n = len(group_data)
|
n = len(group_data)
|
||||||
d_data = {l[0]: (l[1], l[2], l[3]) for l in group_data}
|
d_data = {l[0]: (l[1], l[2], l[3]) for l in group_data}
|
||||||
@@ -466,7 +475,7 @@ def process_single_task(task_tuple, precomputed_response=None):
|
|||||||
for (i,f) in enumerate(res["feedback"]):
|
for (i,f) in enumerate(res["feedback"]):
|
||||||
b = f.get("box_2d")
|
b = f.get("box_2d")
|
||||||
if b:
|
if b:
|
||||||
ymin, xmin, ymax, xmax = b
|
ymin, _xmin, ymax, xmax = b
|
||||||
ymin = ymin * total_height // 1000
|
ymin = ymin * total_height // 1000
|
||||||
ymax = ymax * total_height // 1000
|
ymax = ymax * total_height // 1000
|
||||||
|
|
||||||
@@ -493,7 +502,7 @@ def process_single_task(task_tuple, precomputed_response=None):
|
|||||||
res["feedback"] = correct_boxes_with_gemini(
|
res["feedback"] = correct_boxes_with_gemini(
|
||||||
pid, label, pdf_path, res["feedback"],
|
pid, label, pdf_path, res["feedback"],
|
||||||
yming, ymaxg, width_r, total_height)
|
yming, ymaxg, width_r, total_height)
|
||||||
except Exception as e:
|
except Exception as e: # noqa: BLE001 - correction fallback
|
||||||
tprint(f"\tCorrection failed for Copie {pid}, {group_name} : {e}\n\tRemoving the boxes")
|
tprint(f"\tCorrection failed for Copie {pid}, {group_name} : {e}\n\tRemoving the boxes")
|
||||||
# Fallback if the second request fails entirely
|
# Fallback if the second request fails entirely
|
||||||
for (i, f) in enumerate(res["feedback"]):
|
for (i, f) in enumerate(res["feedback"]):
|
||||||
@@ -516,7 +525,7 @@ def process_single_task(task_tuple, precomputed_response=None):
|
|||||||
tprint(f"Error decoding JSON for {file_path}", file=sys.stderr)
|
tprint(f"Error decoding JSON for {file_path}", file=sys.stderr)
|
||||||
with io_lock:
|
with io_lock:
|
||||||
errors_summary.append(("Error decoding JSON response", file_path))
|
errors_summary.append(("Error decoding JSON response", file_path))
|
||||||
except Exception as e:
|
except Exception as e: # noqa: BLE001 - per-task processing boundary
|
||||||
error_msg = f"Exception processing {file_path}: {e}"
|
error_msg = f"Exception processing {file_path}: {e}"
|
||||||
print(error_msg, file=sys.stderr)
|
print(error_msg, file=sys.stderr)
|
||||||
with io_lock:
|
with io_lock:
|
||||||
@@ -602,19 +611,20 @@ def resolve_delayed_moves():
|
|||||||
|
|
||||||
return new_tasks
|
return new_tasks
|
||||||
|
|
||||||
if __name__ == "__main__":
|
def run_configured(args: argparse.Namespace) -> ExitCode:
|
||||||
|
global client, tasks_to_process
|
||||||
|
if client is None:
|
||||||
|
client = genai.Client(api_key=api_key)
|
||||||
if args.refaire:
|
if args.refaire:
|
||||||
refaire_path = INPUT_DIR / "refaire.json"
|
refaire_path = INPUT_DIR / "refaire.json"
|
||||||
overwritten_path = INPUT_DIR / "overwritten_correction.json"
|
overwritten_path = INPUT_DIR / "overwritten_correction.json"
|
||||||
|
|
||||||
if refaire_path.exists():
|
if refaire_path.exists():
|
||||||
with open(refaire_path, "r", encoding="utf-8") as f:
|
refaire_list = read_json(refaire_path)
|
||||||
refaire_list = json.load(f)
|
|
||||||
|
|
||||||
overwritten_data = []
|
overwritten_data = []
|
||||||
if overwritten_path.exists():
|
if overwritten_path.exists():
|
||||||
with open(overwritten_path, "r", encoding="utf-8") as f:
|
overwritten_data = read_json(overwritten_path)
|
||||||
overwritten_data = json.load(f)
|
|
||||||
|
|
||||||
dirty_results = False
|
dirty_results = False
|
||||||
|
|
||||||
@@ -659,10 +669,12 @@ if __name__ == "__main__":
|
|||||||
pdf_path = copie_dir / f"{label}.pdf"
|
pdf_path = copie_dir / f"{label}.pdf"
|
||||||
is_new = False
|
is_new = False
|
||||||
|
|
||||||
if not pdf_path.exists():
|
if (
|
||||||
if (copie_dir / f"{label}_new.pdf").exists():
|
not pdf_path.exists()
|
||||||
pdf_path = copie_dir / f"{label}_new.pdf"
|
and (copie_dir / f"{label}_new.pdf").exists()
|
||||||
is_new = True
|
):
|
||||||
|
pdf_path = copie_dir / f"{label}_new.pdf"
|
||||||
|
is_new = True
|
||||||
# elif (copie_dir / f"{label}_old.pdf").exists():
|
# elif (copie_dir / f"{label}_old.pdf").exists():
|
||||||
# pdf_path = copie_dir / f"{label}_old.pdf"
|
# pdf_path = copie_dir / f"{label}_old.pdf"
|
||||||
|
|
||||||
@@ -690,7 +702,11 @@ if __name__ == "__main__":
|
|||||||
input(f"About to batch from: {args.batch_from}. Press Enter to confirm...")
|
input(f"About to batch from: {args.batch_from}. Press Enter to confirm...")
|
||||||
break
|
break
|
||||||
if args.batch_from not in all_labels:
|
if args.batch_from not in all_labels:
|
||||||
sys.exit(f"Error: Label '{args.batch_from}' not found. Available labels: {all_labels}")
|
raise CliError(
|
||||||
|
f"Label '{args.batch_from}' not found. Available labels: "
|
||||||
|
f"{all_labels}",
|
||||||
|
ExitCode.INVALID_ARGUMENTS,
|
||||||
|
)
|
||||||
|
|
||||||
target_idx = all_labels.index(args.batch_from)
|
target_idx = all_labels.index(args.batch_from)
|
||||||
live_tasks = []
|
live_tasks = []
|
||||||
@@ -714,58 +730,66 @@ if __name__ == "__main__":
|
|||||||
|
|
||||||
count_flash = 0
|
count_flash = 0
|
||||||
count_pro = 0
|
count_pro = 0
|
||||||
|
flash_lines = []
|
||||||
with open(batch_flash_file, "w", encoding="utf-8") as f_flash, \
|
pro_lines = []
|
||||||
open(batch_pro_file, "w", encoding="utf-8") as f_pro:
|
for task in batch_tasks:
|
||||||
|
file_path, label = task[0], task[1]
|
||||||
for task in batch_tasks:
|
json_path = Path(file_path).with_suffix(".json")
|
||||||
file_path, label = task[0], task[1]
|
group_data = read_json(json_path)
|
||||||
group_name = os.path.splitext(file_path)[0]
|
use_flash = len(group_data) >= 4 or group_data[-1][2] <= 500
|
||||||
json_path = group_name + '.json'
|
b64_img = base64.b64encode(Path(file_path).read_bytes()).decode(
|
||||||
|
"utf-8"
|
||||||
with open(json_path, 'r') as jf:
|
)
|
||||||
group_data = json.load(jf)
|
request = {
|
||||||
use_flash = len(group_data) >= 4 or group_data[-1][2] <= 500
|
"key": file_path,
|
||||||
|
"request": {
|
||||||
image_data = Path(file_path).read_bytes()
|
"contents": [
|
||||||
b64_img = base64.b64encode(image_data).decode("utf-8")
|
{
|
||||||
|
|
||||||
# Format payload matching Gemini Batch API file requirements
|
|
||||||
req = {
|
|
||||||
"key": file_path, # The ID returned in the output file
|
|
||||||
"request": {
|
|
||||||
"contents": [{
|
|
||||||
"role": "user",
|
"role": "user",
|
||||||
"parts": [
|
"parts": [
|
||||||
{"inlineData": {"mimeType": "image/jpeg", "data": b64_img}},
|
{
|
||||||
{"text": prompting.make_prompt(INPUT_DIR,label)}
|
"inlineData": {
|
||||||
]
|
"mimeType": "image/jpeg",
|
||||||
}],
|
"data": b64_img,
|
||||||
"generation_config": {
|
}
|
||||||
"temperature": 1.0,
|
},
|
||||||
"topP": 0.95,
|
{"text": prompting.make_prompt(INPUT_DIR, label)},
|
||||||
"maxOutputTokens": 65535,
|
],
|
||||||
"responseMimeType": "application/json",
|
|
||||||
"responseSchema": prompting.UNROLLED_SCHEMA
|
|
||||||
}
|
}
|
||||||
}
|
],
|
||||||
}
|
"generation_config": {
|
||||||
|
"temperature": 1.0,
|
||||||
|
"topP": 0.95,
|
||||||
|
"maxOutputTokens": 65535,
|
||||||
|
"responseMimeType": "application/json",
|
||||||
|
"responseSchema": prompting.UNROLLED_SCHEMA,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
line = json.dumps(request)
|
||||||
|
if use_flash:
|
||||||
|
flash_lines.append(line)
|
||||||
|
count_flash += 1
|
||||||
|
else:
|
||||||
|
pro_lines.append(line)
|
||||||
|
count_pro += 1
|
||||||
|
atomic_write_text(
|
||||||
|
batch_flash_file,
|
||||||
|
"\n".join(flash_lines) + ("\n" if flash_lines else ""),
|
||||||
|
)
|
||||||
|
atomic_write_text(
|
||||||
|
batch_pro_file,
|
||||||
|
"\n".join(pro_lines) + ("\n" if pro_lines else ""),
|
||||||
|
)
|
||||||
|
|
||||||
if use_flash:
|
print("Batch generation complete.")
|
||||||
f_flash.write(json.dumps(req) + "\n")
|
|
||||||
count_flash += 1
|
|
||||||
else:
|
|
||||||
f_pro.write(json.dumps(req) + "\n")
|
|
||||||
count_pro += 1
|
|
||||||
|
|
||||||
print(f"Batch generation complete.")
|
|
||||||
print(f" - {count_flash} requests saved to {batch_flash_file} (for {MODEL_ID_flash})")
|
print(f" - {count_flash} requests saved to {batch_flash_file} (for {MODEL_ID_flash})")
|
||||||
print(f" - {count_pro} requests saved to {batch_pro_file} (for {MODEL_ID_pro})")
|
print(f" - {count_pro} requests saved to {batch_pro_file} (for {MODEL_ID_pro})")
|
||||||
print("Upload these files via the File API and create two separate batch jobs.")
|
print("Upload these files via the File API and create two separate batch jobs.")
|
||||||
|
|
||||||
# If there's no live tasks to do, and we aren't doing a batched ingestion, exit right away
|
# If there's no live tasks to do, and we aren't doing a batched ingestion, exit right away
|
||||||
if not tasks_to_process and not args.deal_with_batched:
|
if not tasks_to_process and not args.deal_with_batched:
|
||||||
sys.exit(0)
|
return ExitCode.SUCCESS
|
||||||
|
|
||||||
batched_responses = {}
|
batched_responses = {}
|
||||||
if args.deal_with_batched:
|
if args.deal_with_batched:
|
||||||
@@ -807,8 +831,11 @@ if __name__ == "__main__":
|
|||||||
if new_generated_tasks:
|
if new_generated_tasks:
|
||||||
for new_task in new_generated_tasks:
|
for new_task in new_generated_tasks:
|
||||||
futures[executor.submit(process_single_task, new_task)] = new_task
|
futures[executor.submit(process_single_task, new_task)] = new_task
|
||||||
except Exception as e:
|
except Exception as e: # noqa: BLE001 - future boundary
|
||||||
print(f"Exception during task execution: {e}", file=sys.stderr)
|
print(f"Exception during task execution: {e}", file=sys.stderr)
|
||||||
|
failed_task = futures[future]
|
||||||
|
with io_lock:
|
||||||
|
errors_summary.append((str(e), failed_task[0]))
|
||||||
|
|
||||||
tasks_to_process = [] # Vider la liste une fois traitée
|
tasks_to_process = [] # Vider la liste une fois traitée
|
||||||
|
|
||||||
@@ -840,9 +867,12 @@ if __name__ == "__main__":
|
|||||||
|
|
||||||
if unresolved_delayed:
|
if unresolved_delayed:
|
||||||
manual_path = INPUT_DIR / "manual_resolutions.txt"
|
manual_path = INPUT_DIR / "manual_resolutions.txt"
|
||||||
with open(manual_path, "w", encoding="utf-8") as f:
|
atomic_write_text(
|
||||||
f.write("### Use -> x>, -x, ss, sx, xx, xs\n")
|
manual_path,
|
||||||
f.write("\n".join(unresolved_delayed) + "\n")
|
"### Use -> x>, -x, ss, sx, xx, xs\n"
|
||||||
|
+ "\n".join(unresolved_delayed)
|
||||||
|
+ "\n",
|
||||||
|
)
|
||||||
print(f"\n[!] Unresolved delayed tasks found! Wrote to {manual_path}.")
|
print(f"\n[!] Unresolved delayed tasks found! Wrote to {manual_path}.")
|
||||||
print(" Please edit it manually, then run `python resolve_manual.py <InputDir>`")
|
print(" Please edit it manually, then run `python resolve_manual.py <InputDir>`")
|
||||||
|
|
||||||
@@ -855,3 +885,101 @@ if __name__ == "__main__":
|
|||||||
print(err, file=sys.stderr)
|
print(err, file=sys.stderr)
|
||||||
escaped_path = shlex.quote(str(file))
|
escaped_path = shlex.quote(str(file))
|
||||||
print(f"Run : python correction.py {escaped_path}")
|
print(f"Run : python correction.py {escaped_path}")
|
||||||
|
return ExitCode.PARTIAL if errors_summary else ExitCode.SUCCESS
|
||||||
|
|
||||||
|
|
||||||
|
def run(
|
||||||
|
workspace: EvaluationWorkspace,
|
||||||
|
targets: list[Path],
|
||||||
|
args: argparse.Namespace,
|
||||||
|
*,
|
||||||
|
api_client=None,
|
||||||
|
) -> ExitCode:
|
||||||
|
if args.reset:
|
||||||
|
workspace.require_directories("Copies")
|
||||||
|
reset_workspace(workspace)
|
||||||
|
return ExitCode.SUCCESS
|
||||||
|
workspace.require_directories("Copies", "Par label")
|
||||||
|
workspace.require_files("labels")
|
||||||
|
if args.refaire:
|
||||||
|
workspace.require_files("refaire.json")
|
||||||
|
discovered, warnings = discover_tasks(workspace, targets)
|
||||||
|
for warning in warnings:
|
||||||
|
print(f"Warning: {warning}")
|
||||||
|
configure_runtime(workspace, discovered, args, api_client=api_client)
|
||||||
|
if not discovered and not args.refaire:
|
||||||
|
return ExitCode.PARTIAL
|
||||||
|
try:
|
||||||
|
status = run_configured(args)
|
||||||
|
finally:
|
||||||
|
for thread_id in list(thread_logs):
|
||||||
|
flush_thread_log(thread_id)
|
||||||
|
if warnings and status == ExitCode.SUCCESS:
|
||||||
|
return ExitCode.PARTIAL
|
||||||
|
return status
|
||||||
|
|
||||||
|
|
||||||
|
def build_parser() -> argparse.ArgumentParser:
|
||||||
|
parser = target_parser("Correct grouped answers with Gemini")
|
||||||
|
parser.add_argument(
|
||||||
|
"additional_targets",
|
||||||
|
nargs="*",
|
||||||
|
type=Path,
|
||||||
|
help="Additional group JPG files from the same evaluation",
|
||||||
|
)
|
||||||
|
parser.add_argument("--overwrite", action="store_true", help="Redo requests")
|
||||||
|
parser.add_argument("--limit", type=int, help="Maximum Gemini Pro calls")
|
||||||
|
parser.add_argument(
|
||||||
|
"--refaire",
|
||||||
|
action="store_true",
|
||||||
|
help="Redo copies and labels listed in refaire.json",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--batch",
|
||||||
|
action="store_true",
|
||||||
|
help="Generate Gemini batch request JSONL files",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--batch-from",
|
||||||
|
metavar="LABEL",
|
||||||
|
help="Process earlier labels live and batch from LABEL onward",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--deal-with-batched",
|
||||||
|
action="store_true",
|
||||||
|
help="Consume batched_correction_result.jsonl",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--reset",
|
||||||
|
action="store_true",
|
||||||
|
help="Delete correction state, restore _old PDFs, and delete _new PDFs",
|
||||||
|
)
|
||||||
|
return parser
|
||||||
|
|
||||||
|
|
||||||
|
def main(argv: Sequence[str] | None = None) -> int:
|
||||||
|
parser = build_parser()
|
||||||
|
|
||||||
|
def handle(args: argparse.Namespace) -> ExitCode:
|
||||||
|
workspace, target = workspace_from_target(args)
|
||||||
|
targets = [target]
|
||||||
|
for additional in args.additional_targets:
|
||||||
|
resolved = additional.expanduser().resolve()
|
||||||
|
if not resolved.exists():
|
||||||
|
raise CliError(
|
||||||
|
f"Target does not exist: {resolved}",
|
||||||
|
ExitCode.INVALID_WORKSPACE,
|
||||||
|
)
|
||||||
|
if EvaluationWorkspace.discover(resolved).root != workspace.root:
|
||||||
|
raise CliError(
|
||||||
|
"All targets must belong to the same evaluation",
|
||||||
|
ExitCode.INVALID_ARGUMENTS,
|
||||||
|
)
|
||||||
|
targets.append(resolved)
|
||||||
|
return run(workspace, targets, args)
|
||||||
|
|
||||||
|
return execute(parser, argv, handle)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
raise SystemExit(main())
|
||||||
|
|||||||
+70
-50
@@ -1,63 +1,83 @@
|
|||||||
import os
|
from __future__ import annotations
|
||||||
import sys
|
|
||||||
import argparse
|
import argparse
|
||||||
from pathlib import Path
|
from collections.abc import Sequence
|
||||||
|
|
||||||
from google import genai
|
from google import genai
|
||||||
|
|
||||||
def main():
|
import config
|
||||||
parser = argparse.ArgumentParser(description="Download and combine completed batch jobs for a directory.")
|
from copienator import (
|
||||||
parser.add_argument("root_dir", type=str, help="Directory containing the original batches")
|
CliError,
|
||||||
args = parser.parse_args()
|
EvaluationWorkspace,
|
||||||
|
ExitCode,
|
||||||
|
atomic_write_bytes,
|
||||||
|
evaluation_parser,
|
||||||
|
execute,
|
||||||
|
read_json,
|
||||||
|
workspace_from_args,
|
||||||
|
)
|
||||||
|
|
||||||
target_dir = Path(args.root_dir)
|
|
||||||
dir_name = target_dir.name
|
|
||||||
output_path = target_dir / "batched_correction_result.jsonl"
|
|
||||||
|
|
||||||
if "GEMINI_API_KEY" not in os.environ:
|
def run(workspace: EvaluationWorkspace, *, client=None) -> ExitCode:
|
||||||
sys.exit("Error: GEMINI_API_KEY environment variable not set.")
|
if client is None:
|
||||||
|
if not config.API_KEY:
|
||||||
|
raise CliError("GEMINI_API_KEY is not configured")
|
||||||
|
client = genai.Client(api_key=config.API_KEY)
|
||||||
|
matching = []
|
||||||
|
if workspace.batch_jobs_file.is_file():
|
||||||
|
manifest = read_json(workspace.batch_jobs_file)
|
||||||
|
jobs = manifest.get("jobs") if isinstance(manifest, dict) else None
|
||||||
|
if not isinstance(jobs, dict):
|
||||||
|
raise CliError(f"Invalid batch manifest: {workspace.batch_jobs_file}")
|
||||||
|
matching = [
|
||||||
|
client.batches.get(name=entry["name"])
|
||||||
|
for entry in jobs.values()
|
||||||
|
if isinstance(entry, dict) and isinstance(entry.get("name"), str)
|
||||||
|
]
|
||||||
|
else:
|
||||||
|
matching = [
|
||||||
|
job
|
||||||
|
for job in client.batches.list()
|
||||||
|
if workspace.name in str(getattr(job, "display_name", ""))
|
||||||
|
]
|
||||||
|
if not matching:
|
||||||
|
raise CliError(
|
||||||
|
f"No batch jobs found for evaluation {workspace.name!r}"
|
||||||
|
)
|
||||||
|
for job in matching:
|
||||||
|
state = job.state.name if hasattr(job.state, "name") else job.state
|
||||||
|
print(f"{job.display_name}: {state}")
|
||||||
|
if state != "JOB_STATE_SUCCEEDED":
|
||||||
|
print("Not all matching jobs have succeeded yet.")
|
||||||
|
return ExitCode.PARTIAL
|
||||||
|
|
||||||
client = genai.Client()
|
chunks = []
|
||||||
|
incomplete = False
|
||||||
|
for job in matching:
|
||||||
|
destination = getattr(job, "dest", None)
|
||||||
|
file_name = getattr(destination, "file_name", None)
|
||||||
|
if not file_name:
|
||||||
|
print(f"Warning: {job.display_name} has no output file.")
|
||||||
|
incomplete = True
|
||||||
|
continue
|
||||||
|
payload = client.files.download(file=file_name)
|
||||||
|
chunks.append(payload.rstrip(b"\n"))
|
||||||
|
if not chunks:
|
||||||
|
return ExitCode.PARTIAL
|
||||||
|
output_path = workspace.batched_correction_result_file
|
||||||
|
atomic_write_bytes(output_path, b"\n".join(chunks) + b"\n")
|
||||||
|
print(f"Saved combined batch results to {output_path}")
|
||||||
|
return ExitCode.PARTIAL if incomplete else ExitCode.SUCCESS
|
||||||
|
|
||||||
print(f"Fetching jobs matching '{dir_name}'...")
|
|
||||||
all_jobs = client.batches.list()
|
|
||||||
matching_jobs = []
|
|
||||||
|
|
||||||
# 1. Find jobs associated with this directory
|
def build_parser() -> argparse.ArgumentParser:
|
||||||
for job in all_jobs:
|
return evaluation_parser("Download and combine correction batch results")
|
||||||
if hasattr(job, 'display_name') and job.display_name and dir_name in job.display_name:
|
|
||||||
matching_jobs.append(job)
|
|
||||||
|
|
||||||
if not matching_jobs:
|
|
||||||
sys.exit(f"No batch jobs found containing '{dir_name}' in their display name.")
|
|
||||||
|
|
||||||
# 2. Check that all matching jobs are complete
|
def main(argv: Sequence[str] | None = None) -> int:
|
||||||
for job in matching_jobs:
|
parser = build_parser()
|
||||||
state = job.state.name if hasattr(job.state, 'name') else job.state
|
return execute(parser, argv, lambda args: run(workspace_from_args(args)))
|
||||||
print(f"Found Job: {job.display_name} | State: {state}")
|
|
||||||
if state != 'JOB_STATE_SUCCEEDED':
|
|
||||||
sys.exit(f"Error: Job '{job.display_name}' has not succeeded yet. Try again later.")
|
|
||||||
|
|
||||||
# 3. Download and concatenate
|
|
||||||
print("\nAll jobs succeeded. Downloading results...")
|
|
||||||
combined_data = b""
|
|
||||||
|
|
||||||
for job in matching_jobs:
|
|
||||||
if hasattr(job, 'dest') and job.dest and hasattr(job.dest, 'file_name') and job.dest.file_name:
|
|
||||||
print(f"Downloading output for {job.display_name}...")
|
|
||||||
file_content_bytes = client.files.download(file=job.dest.file_name)
|
|
||||||
|
|
||||||
combined_data += file_content_bytes
|
|
||||||
# Ensure proper line separation between files in JSONL
|
|
||||||
if combined_data and not combined_data.endswith(b'\n'):
|
|
||||||
combined_data += b'\n'
|
|
||||||
else:
|
|
||||||
print(f"Warning: Job {job.display_name} succeeded but has no output file.")
|
|
||||||
|
|
||||||
# 4. Save to destination
|
|
||||||
with open(output_path, "wb") as f:
|
|
||||||
f.write(combined_data)
|
|
||||||
|
|
||||||
print(f"\nSuccess! All results concatenated and saved to:\n{output_path}")
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
main()
|
raise SystemExit(main())
|
||||||
|
|||||||
+76
-57
@@ -1,77 +1,96 @@
|
|||||||
import os
|
from __future__ import annotations
|
||||||
import sys
|
|
||||||
import argparse
|
import argparse
|
||||||
from pathlib import Path
|
from collections.abc import Sequence
|
||||||
|
|
||||||
from google import genai
|
from google import genai
|
||||||
from google.genai import types
|
from google.genai import types
|
||||||
|
|
||||||
def main():
|
import config
|
||||||
parser = argparse.ArgumentParser(description="Upload JSONL files and create Gemini Batch jobs.")
|
from copienator import (
|
||||||
parser.add_argument("root_dir", type=str, help="Root directory containing the batch JSONL files")
|
CliError,
|
||||||
args = parser.parse_args()
|
EvaluationWorkspace,
|
||||||
|
ExitCode,
|
||||||
|
atomic_write_json,
|
||||||
|
evaluation_parser,
|
||||||
|
execute,
|
||||||
|
read_json,
|
||||||
|
workspace_from_args,
|
||||||
|
)
|
||||||
|
|
||||||
root_dir = Path(args.root_dir)
|
|
||||||
|
|
||||||
if "GEMINI_API_KEY" not in os.environ:
|
def run(workspace: EvaluationWorkspace, *, client=None) -> ExitCode:
|
||||||
sys.exit("Error: GEMINI_API_KEY environment variable not set.")
|
if client is None:
|
||||||
|
if not config.API_KEY:
|
||||||
client = genai.Client()
|
raise CliError("GEMINI_API_KEY is not configured")
|
||||||
|
client = genai.Client(api_key=config.API_KEY)
|
||||||
# Define the batch files and their corresponding models
|
batches = (
|
||||||
batches_to_create = [
|
(
|
||||||
{
|
"flash",
|
||||||
"file_path": root_dir / "batch_requests_flash.jsonl",
|
workspace.root / "batch_requests_flash.jsonl",
|
||||||
"model_id": "gemini-3-flash-preview",
|
config.MODEL_FLASH_ID,
|
||||||
"display_name": f"flash-correction-{root_dir.name}"
|
f"flash-correction-{workspace.name}",
|
||||||
},
|
),
|
||||||
{
|
(
|
||||||
"file_path": root_dir / "batch_requests_pro.jsonl",
|
"pro",
|
||||||
"model_id": "gemini-3.1-pro-preview",
|
workspace.root / "batch_requests_pro.jsonl",
|
||||||
"display_name": f"pro-correction-{root_dir.name}"
|
config.MODEL_PRO_ID,
|
||||||
}
|
f"pro-correction-{workspace.name}",
|
||||||
]
|
),
|
||||||
|
)
|
||||||
for batch in batches_to_create:
|
manifest = {
|
||||||
file_path = batch["file_path"]
|
"version": 1,
|
||||||
model_id = batch["model_id"]
|
"evaluation": workspace.name,
|
||||||
display_name = batch["display_name"]
|
"jobs": {},
|
||||||
|
}
|
||||||
# Check if the file exists
|
if workspace.batch_jobs_file.is_file():
|
||||||
if not file_path.exists():
|
previous = read_json(workspace.batch_jobs_file)
|
||||||
|
if isinstance(previous, dict) and isinstance(previous.get("jobs"), dict):
|
||||||
|
manifest["jobs"] = previous["jobs"]
|
||||||
|
started = 0
|
||||||
|
for tier, file_path, model_id, display_name in batches:
|
||||||
|
if not file_path.is_file():
|
||||||
print(f"Skipping {model_id}: {file_path.name} does not exist.")
|
print(f"Skipping {model_id}: {file_path.name} does not exist.")
|
||||||
continue
|
continue
|
||||||
|
|
||||||
# Check if the file is empty (e.g., if all tasks went to Flash, Pro might be empty)
|
|
||||||
if file_path.stat().st_size == 0:
|
if file_path.stat().st_size == 0:
|
||||||
print(f"Skipping {model_id}: {file_path.name} is empty.")
|
print(f"Skipping {model_id}: {file_path.name} is empty.")
|
||||||
continue
|
continue
|
||||||
|
print(f"Uploading {file_path.name} for model {model_id}...")
|
||||||
print(f"Processing {file_path.name} for model {model_id}...")
|
uploaded = client.files.upload(
|
||||||
|
|
||||||
# 1. Upload the file to the File API
|
|
||||||
print(f" Uploading file...")
|
|
||||||
uploaded_file = client.files.upload(
|
|
||||||
file=str(file_path),
|
file=str(file_path),
|
||||||
config=types.UploadFileConfig(
|
config=types.UploadFileConfig(
|
||||||
display_name=f"{display_name}-input",
|
display_name=f"{display_name}-input",
|
||||||
mime_type='jsonl'
|
mime_type="jsonl",
|
||||||
)
|
),
|
||||||
)
|
)
|
||||||
print(f" Uploaded successfully! File ID: {uploaded_file.name}")
|
job = client.batches.create(
|
||||||
|
|
||||||
# 2. Create the batch job
|
|
||||||
print(f" Starting batch job...")
|
|
||||||
batch_job = client.batches.create(
|
|
||||||
model=model_id,
|
model=model_id,
|
||||||
src=uploaded_file.name,
|
src=uploaded.name,
|
||||||
config={
|
config={"display_name": display_name},
|
||||||
'display_name': display_name,
|
|
||||||
},
|
|
||||||
)
|
)
|
||||||
print(f" Success! Batch Job Name: {batch_job.name}\n")
|
started += 1
|
||||||
|
manifest["jobs"][tier] = {
|
||||||
|
"name": job.name,
|
||||||
|
"display_name": display_name,
|
||||||
|
"model": model_id,
|
||||||
|
"request_file": file_path.name,
|
||||||
|
}
|
||||||
|
atomic_write_json(workspace.batch_jobs_file, manifest)
|
||||||
|
print(f"Started batch job: {job.name}")
|
||||||
|
if not started:
|
||||||
|
print("No non-empty batch request files were found.")
|
||||||
|
return ExitCode.PARTIAL
|
||||||
|
return ExitCode.SUCCESS
|
||||||
|
|
||||||
|
|
||||||
|
def build_parser() -> argparse.ArgumentParser:
|
||||||
|
return evaluation_parser("Upload correction JSONL files and start Gemini batches")
|
||||||
|
|
||||||
|
|
||||||
|
def main(argv: Sequence[str] | None = None) -> int:
|
||||||
|
parser = build_parser()
|
||||||
|
return execute(parser, argv, lambda args: run(workspace_from_args(args)))
|
||||||
|
|
||||||
print("-" * 50)
|
|
||||||
print("All batch jobs have been initiated.")
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
main()
|
raise SystemExit(main())
|
||||||
|
|||||||
@@ -2,6 +2,7 @@ from __future__ import annotations
|
|||||||
|
|
||||||
import importlib.util
|
import importlib.util
|
||||||
import io
|
import io
|
||||||
|
import json
|
||||||
import os
|
import os
|
||||||
import queue
|
import queue
|
||||||
import sys
|
import sys
|
||||||
@@ -11,6 +12,7 @@ import unittest
|
|||||||
from concurrent.futures import ThreadPoolExecutor
|
from concurrent.futures import ThreadPoolExecutor
|
||||||
from contextlib import redirect_stderr
|
from contextlib import redirect_stderr
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
from types import SimpleNamespace
|
||||||
from unittest.mock import Mock, patch
|
from unittest.mock import Mock, patch
|
||||||
|
|
||||||
from PIL import Image
|
from PIL import Image
|
||||||
@@ -21,6 +23,7 @@ from copienator import (
|
|||||||
WorkspaceNotFoundError,
|
WorkspaceNotFoundError,
|
||||||
WorkspaceValidationError,
|
WorkspaceValidationError,
|
||||||
atomic_update_json,
|
atomic_update_json,
|
||||||
|
atomic_write_bytes,
|
||||||
atomic_write_json,
|
atomic_write_json,
|
||||||
read_json,
|
read_json,
|
||||||
workspace_from_target,
|
workspace_from_target,
|
||||||
@@ -114,6 +117,12 @@ class WorkspaceTests(unittest.TestCase):
|
|||||||
|
|
||||||
|
|
||||||
class AtomicJsonTests(unittest.TestCase):
|
class AtomicJsonTests(unittest.TestCase):
|
||||||
|
def test_atomic_binary_round_trip(self) -> None:
|
||||||
|
with tempfile.TemporaryDirectory() as directory:
|
||||||
|
path = Path(directory) / "result.jsonl"
|
||||||
|
atomic_write_bytes(path, b'{"one":1}\n')
|
||||||
|
self.assertEqual(path.read_bytes(), b'{"one":1}\n')
|
||||||
|
|
||||||
def test_atomic_round_trip_and_unicode(self) -> None:
|
def test_atomic_round_trip_and_unicode(self) -> None:
|
||||||
with tempfile.TemporaryDirectory() as directory:
|
with tempfile.TemporaryDirectory() as directory:
|
||||||
path = Path(directory) / "nested" / "state.json"
|
path = Path(directory) / "nested" / "state.json"
|
||||||
@@ -313,6 +322,14 @@ class StandardCliTests(unittest.TestCase):
|
|||||||
"gemini_for_labels": load_script_module(
|
"gemini_for_labels": load_script_module(
|
||||||
"gemini_for_labels.py", "gemini_for_labels"
|
"gemini_for_labels.py", "gemini_for_labels"
|
||||||
),
|
),
|
||||||
|
"correction": load_script_module("correction.py", "correction"),
|
||||||
|
"submit_batches": load_script_module(
|
||||||
|
"submit_batches.py", "submit_batches"
|
||||||
|
),
|
||||||
|
"batch_status": load_script_module("batch_status.py", "batch_status"),
|
||||||
|
"fetch_batched_results": load_script_module(
|
||||||
|
"fetch_batched_results.py", "fetch_batched_results"
|
||||||
|
),
|
||||||
"copies_tools": load_script_module(
|
"copies_tools": load_script_module(
|
||||||
"copies_tools.py", "copienator_copies_tools_test"
|
"copies_tools.py", "copienator_copies_tools_test"
|
||||||
),
|
),
|
||||||
@@ -355,6 +372,9 @@ class StandardCliTests(unittest.TestCase):
|
|||||||
"page_splitter": [missing],
|
"page_splitter": [missing],
|
||||||
"plotting": [missing],
|
"plotting": [missing],
|
||||||
"gemini_for_labels": [missing],
|
"gemini_for_labels": [missing],
|
||||||
|
"correction": [missing],
|
||||||
|
"submit_batches": [missing],
|
||||||
|
"fetch_batched_results": [missing],
|
||||||
}
|
}
|
||||||
for name, arguments in invocations.items():
|
for name, arguments in invocations.items():
|
||||||
with self.subTest(script=name), redirect_stderr(io.StringIO()):
|
with self.subTest(script=name), redirect_stderr(io.StringIO()):
|
||||||
@@ -472,6 +492,21 @@ class StandardCliTests(unittest.TestCase):
|
|||||||
"default",
|
"default",
|
||||||
{"target": evaluation, "overwrite": True},
|
{"target": evaluation, "overwrite": True},
|
||||||
),
|
),
|
||||||
|
"correction": (
|
||||||
|
"correction",
|
||||||
|
"live",
|
||||||
|
{"target": evaluation, "overwrite": True, "limit": 5},
|
||||||
|
),
|
||||||
|
"submit_batches": (
|
||||||
|
"submit_batches",
|
||||||
|
"default",
|
||||||
|
{"target": evaluation},
|
||||||
|
),
|
||||||
|
"fetch_batched_results": (
|
||||||
|
"fetch_batches",
|
||||||
|
"default",
|
||||||
|
{"target": evaluation},
|
||||||
|
),
|
||||||
}
|
}
|
||||||
for module_name, (step_id, variant_id, values) in cases.items():
|
for module_name, (step_id, variant_id, values) in cases.items():
|
||||||
step = steps[step_id]
|
step = steps[step_id]
|
||||||
@@ -900,6 +935,135 @@ class StandardCliTests(unittest.TestCase):
|
|||||||
"Ex 1",
|
"Ex 1",
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def test_correction_overwrite_keeps_previous_state_until_a_commit(self) -> None:
|
||||||
|
module = self.modules["correction"]
|
||||||
|
with tempfile.TemporaryDirectory() as directory:
|
||||||
|
evaluation = Path(directory) / "Exam"
|
||||||
|
evaluation.mkdir()
|
||||||
|
correction = evaluation / "correction.json"
|
||||||
|
progress = evaluation / "correction_progress.json"
|
||||||
|
atomic_write_json(correction, {"Ex 1": [[{"id": "01"}]]})
|
||||||
|
atomic_write_json(progress, [["old.jpg", "Ex 1"]])
|
||||||
|
args = module.build_parser().parse_args(
|
||||||
|
[str(evaluation), "--overwrite"]
|
||||||
|
)
|
||||||
|
module.configure_runtime(
|
||||||
|
EvaluationWorkspace(evaluation),
|
||||||
|
[("new.jpg", "Ex 1")],
|
||||||
|
args,
|
||||||
|
api_client=Mock(),
|
||||||
|
)
|
||||||
|
self.assertEqual(read_json(correction), {"Ex 1": [[{"id": "01"}]]})
|
||||||
|
self.assertEqual(read_json(progress), [["old.jpg", "Ex 1"]])
|
||||||
|
self.assertEqual(module.results, {"Ex 1": []})
|
||||||
|
self.assertEqual(module.tasks_to_process, [("new.jpg", "Ex 1")])
|
||||||
|
|
||||||
|
def test_correction_reset_restores_old_and_deletes_new_files(self) -> None:
|
||||||
|
module = self.modules["correction"]
|
||||||
|
with tempfile.TemporaryDirectory() as directory:
|
||||||
|
evaluation = Path(directory) / "Exam"
|
||||||
|
copy_dir = evaluation / "Copies" / "Copie01"
|
||||||
|
copy_dir.mkdir(parents=True)
|
||||||
|
atomic_write_json(evaluation / "correction.json", {"old": True})
|
||||||
|
atomic_write_json(evaluation / "correction_progress.json", ["old"])
|
||||||
|
(copy_dir / "Ex 1.pdf").write_bytes(b"current")
|
||||||
|
(copy_dir / "Ex 1_old.pdf").write_bytes(b"original")
|
||||||
|
(copy_dir / "Ex 2_new.pdf").write_bytes(b"generated")
|
||||||
|
|
||||||
|
self.assertEqual(module.main([str(evaluation), "--reset"]), 0)
|
||||||
|
self.assertFalse((evaluation / "correction.json").exists())
|
||||||
|
self.assertFalse((evaluation / "correction_progress.json").exists())
|
||||||
|
self.assertEqual((copy_dir / "Ex 1.pdf").read_bytes(), b"original")
|
||||||
|
self.assertFalse((copy_dir / "Ex 1_old.pdf").exists())
|
||||||
|
self.assertFalse((copy_dir / "Ex 2_new.pdf").exists())
|
||||||
|
|
||||||
|
def test_correction_batch_request_files_are_written_atomically(self) -> None:
|
||||||
|
module = self.modules["correction"]
|
||||||
|
with tempfile.TemporaryDirectory() as directory:
|
||||||
|
evaluation = Path(directory) / "Exam"
|
||||||
|
group_dir = evaluation / "Par label" / "Ex 1"
|
||||||
|
(evaluation / "Copies").mkdir(parents=True)
|
||||||
|
group_dir.mkdir(parents=True)
|
||||||
|
(evaluation / "labels").write_text("Ex 1\n", encoding="utf-8")
|
||||||
|
image = group_dir / "Group_1.jpg"
|
||||||
|
image.write_bytes(b"image")
|
||||||
|
atomic_write_json(
|
||||||
|
image.with_suffix(".json"),
|
||||||
|
[["01", 0, 400, 1.0, "Ex 1"]],
|
||||||
|
)
|
||||||
|
args = module.build_parser().parse_args([str(evaluation), "--batch"])
|
||||||
|
module.configure_runtime(
|
||||||
|
EvaluationWorkspace(evaluation),
|
||||||
|
[(str(image), "Ex 1")],
|
||||||
|
args,
|
||||||
|
api_client=Mock(),
|
||||||
|
)
|
||||||
|
with patch.object(module.prompting, "make_prompt", return_value="prompt"):
|
||||||
|
self.assertEqual(module.run_configured(args), 0)
|
||||||
|
lines = (evaluation / "batch_requests_flash.jsonl").read_text().splitlines()
|
||||||
|
self.assertEqual(len(lines), 1)
|
||||||
|
self.assertEqual(json.loads(lines[0])["key"], str(image))
|
||||||
|
self.assertEqual(
|
||||||
|
(evaluation / "batch_requests_pro.jsonl").read_text(), ""
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_batch_helpers_use_mocked_api_and_atomic_combination(self) -> None:
|
||||||
|
submit = self.modules["submit_batches"]
|
||||||
|
fetch = self.modules["fetch_batched_results"]
|
||||||
|
status = self.modules["batch_status"]
|
||||||
|
with tempfile.TemporaryDirectory() as directory:
|
||||||
|
evaluation = Path(directory) / "Exam"
|
||||||
|
evaluation.mkdir()
|
||||||
|
(evaluation / "batch_requests_flash.jsonl").write_text(
|
||||||
|
"{}\n", encoding="utf-8"
|
||||||
|
)
|
||||||
|
client = Mock()
|
||||||
|
client.files.upload.return_value = SimpleNamespace(name="files/input")
|
||||||
|
client.batches.create.return_value = SimpleNamespace(name="batches/1")
|
||||||
|
self.assertEqual(
|
||||||
|
submit.run(EvaluationWorkspace(evaluation), client=client), 0
|
||||||
|
)
|
||||||
|
self.assertEqual(client.batches.create.call_count, 1)
|
||||||
|
manifest = json.loads((evaluation / "batch_jobs.json").read_text())
|
||||||
|
self.assertEqual(manifest["jobs"]["flash"]["name"], "batches/1")
|
||||||
|
|
||||||
|
jobs = [
|
||||||
|
SimpleNamespace(
|
||||||
|
name="batches/1",
|
||||||
|
display_name=f"flash-correction-{evaluation.name}",
|
||||||
|
state=SimpleNamespace(name="JOB_STATE_SUCCEEDED"),
|
||||||
|
dest=SimpleNamespace(file_name="files/flash-result"),
|
||||||
|
),
|
||||||
|
SimpleNamespace(
|
||||||
|
name="batches/2",
|
||||||
|
display_name=f"pro-correction-{evaluation.name}",
|
||||||
|
state=SimpleNamespace(name="JOB_STATE_SUCCEEDED"),
|
||||||
|
dest=SimpleNamespace(file_name="files/pro-result"),
|
||||||
|
),
|
||||||
|
]
|
||||||
|
client.batches.get.return_value = jobs[0]
|
||||||
|
client.files.download.side_effect = [b'{"flash":1}\n', b'{"pro":1}']
|
||||||
|
self.assertEqual(
|
||||||
|
fetch.run(EvaluationWorkspace(evaluation), client=client), 0
|
||||||
|
)
|
||||||
|
self.assertEqual(
|
||||||
|
(evaluation / "batched_correction_result.jsonl").read_bytes(),
|
||||||
|
b'{"flash":1}\n',
|
||||||
|
)
|
||||||
|
|
||||||
|
job = SimpleNamespace(
|
||||||
|
state=SimpleNamespace(name="JOB_STATE_SUCCEEDED"),
|
||||||
|
dest=SimpleNamespace(file_name="files/result"),
|
||||||
|
)
|
||||||
|
client.batches.get.return_value = job
|
||||||
|
client.files.download.side_effect = None
|
||||||
|
client.files.download.return_value = b'{"downloaded":true}\n'
|
||||||
|
output = evaluation / "one-result.jsonl"
|
||||||
|
self.assertEqual(
|
||||||
|
status.download_job("batches/1", output=output, client=client), 0
|
||||||
|
)
|
||||||
|
self.assertEqual(output.read_bytes(), b'{"downloaded":true}\n')
|
||||||
|
|
||||||
def test_post_correction_main_cleans_json_atomically(self) -> None:
|
def test_post_correction_main_cleans_json_atomically(self) -> None:
|
||||||
module = self.modules["post_correction"]
|
module = self.modules["post_correction"]
|
||||||
with tempfile.TemporaryDirectory() as directory:
|
with tempfile.TemporaryDirectory() as directory:
|
||||||
|
|||||||
Reference in New Issue
Block a user