fix: make --resume in batch pipeline actually skip completed work
Previously --resume re-queried every gradient in every row, overwriting existing values — an interrupted run could not be resumed cheaply. - bicorder_query.py: add --resume flag; skip gradients whose cells already have values, and report how many were skipped - bicorder_batch.py: pass --resume through to query; skip fully-complete rows before invoking the query script; report partial rows - bicorder_batch.py: import row/config helpers from bicorder_query instead of calling undefined names (would have crashed on --resume)
This commit is contained in:
1 parent
c571bf1c01
commit
fb3bebcea0
2 files changed
+50
-5
No files matched your search
@@ -16,6 +16,9 @@ import argparse
|
|||||||
import subprocess
|
import subprocess
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
|
sys.path.insert(0, str(Path(__file__).parent))
|
||||||
|
from bicorder_query import get_row_values, load_bicorder_config, extract_gradients
|
||||||
|
|
||||||
|
|
||||||
def count_csv_rows(csv_path):
|
def count_csv_rows(csv_path):
|
||||||
"""Count the number of data rows in a CSV file."""
|
"""Count the number of data rows in a CSV file."""
|
||||||
@@ -44,13 +47,15 @@ def run_bicorder_analyze(input_csv, output_csv, bicorder_path, analyst=None, sta
|
|||||||
return True
|
return True
|
||||||
|
|
||||||
|
|
||||||
def query_gradients(output_csv, row_num, bicorder_path, model=None):
|
def query_gradients(output_csv, row_num, bicorder_path, model=None, resume=False):
|
||||||
"""Query all gradients for a protocol row."""
|
"""Query all gradients for a protocol row."""
|
||||||
cmd = ['python3', str(Path(__file__).parent / 'bicorder_query.py'), output_csv, str(row_num),
|
cmd = ['python3', str(Path(__file__).parent / 'bicorder_query.py'), output_csv, str(row_num),
|
||||||
'-b', bicorder_path]
|
'-b', bicorder_path]
|
||||||
|
|
||||||
if model:
|
if model:
|
||||||
cmd.extend(['-m', model])
|
cmd.extend(['-m', model])
|
||||||
|
if resume:
|
||||||
|
cmd.extend(['--resume'])
|
||||||
|
|
||||||
print(f"Starting gradient queries...")
|
print(f"Starting gradient queries...")
|
||||||
|
|
||||||
@@ -64,14 +69,28 @@ def query_gradients(output_csv, row_num, bicorder_path, model=None):
|
|||||||
return True
|
return True
|
||||||
|
|
||||||
|
|
||||||
def process_protocol_row(input_csv, output_csv, row_num, total_rows, bicorder_path, model=None):
|
def process_protocol_row(input_csv, output_csv, row_num, total_rows, bicorder_path, model=None, resume=False):
|
||||||
"""Process a single protocol row through the complete workflow."""
|
"""Process a single protocol row through the complete workflow."""
|
||||||
print(f"\n{'='*60}")
|
print(f"\n{'='*60}")
|
||||||
print(f"Row {row_num}/{total_rows}")
|
print(f"Row {row_num}/{total_rows}")
|
||||||
print(f"{'='*60}")
|
print(f"{'='*60}")
|
||||||
|
|
||||||
|
# With resume: skip rows where all gradient columns already have values
|
||||||
|
if resume:
|
||||||
|
row_values = get_row_values(output_csv, row_num)
|
||||||
|
if row_values:
|
||||||
|
bicorder_data = load_bicorder_config(bicorder_path)
|
||||||
|
gradients = extract_gradients(bicorder_data)
|
||||||
|
gradient_cols = [g['column_name'] for g in gradients]
|
||||||
|
filled = sum(1 for c in gradient_cols if row_values.get(c, '').strip())
|
||||||
|
if filled == len(gradient_cols):
|
||||||
|
print(f"[SKIP] Row {row_num} complete ({filled}/{len(gradient_cols)} values) — resuming")
|
||||||
|
return True
|
||||||
|
elif filled > 0:
|
||||||
|
print(f"[RESUME] Row {row_num} partially complete ({filled}/{len(gradient_cols)} values)")
|
||||||
|
|
||||||
# Query all gradients (each gradient gets a new chat)
|
# Query all gradients (each gradient gets a new chat)
|
||||||
if not query_gradients(output_csv, row_num, bicorder_path, model):
|
if not query_gradients(output_csv, row_num, bicorder_path, model, resume):
|
||||||
print(f"[FAILED] Could not query gradients")
|
print(f"[FAILED] Could not query gradients")
|
||||||
return False
|
return False
|
||||||
|
|
||||||
@@ -112,7 +131,7 @@ Example usage:
|
|||||||
parser.add_argument('--end', type=int,
|
parser.add_argument('--end', type=int,
|
||||||
help='End row number (1-indexed, default: all rows)')
|
help='End row number (1-indexed, default: all rows)')
|
||||||
parser.add_argument('--resume', action='store_true',
|
parser.add_argument('--resume', action='store_true',
|
||||||
help='Resume from existing output CSV (skip rows with values)')
|
help='Resume from existing output CSV (skip gradients that already have values)')
|
||||||
|
|
||||||
args = parser.parse_args()
|
args = parser.parse_args()
|
||||||
|
|
||||||
@@ -156,7 +175,7 @@ Example usage:
|
|||||||
|
|
||||||
for row_num in range(args.start, end_row + 1):
|
for row_num in range(args.start, end_row + 1):
|
||||||
if process_protocol_row(args.input_csv, args.output, row_num, end_row,
|
if process_protocol_row(args.input_csv, args.output, row_num, end_row,
|
||||||
args.bicorder, args.model):
|
args.bicorder, args.model, args.resume):
|
||||||
success_count += 1
|
success_count += 1
|
||||||
else:
|
else:
|
||||||
fail_count += 1
|
fail_count += 1
|
||||||
|
|||||||
@@ -54,6 +54,16 @@ def get_protocol_by_row(csv_path, row_number):
|
|||||||
return None
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def get_row_values(csv_path, row_number):
|
||||||
|
"""Get all existing values for a row (1-indexed) as a dict of column -> value."""
|
||||||
|
with open(csv_path, 'r', encoding='utf-8') as f:
|
||||||
|
reader = csv.DictReader(f)
|
||||||
|
for i, row in enumerate(reader, start=1):
|
||||||
|
if i == row_number:
|
||||||
|
return row
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
def generate_gradient_prompt(protocol_descriptor, protocol_description, gradient):
|
def generate_gradient_prompt(protocol_descriptor, protocol_description, gradient):
|
||||||
"""Generate a prompt for a single gradient evaluation."""
|
"""Generate a prompt for a single gradient evaluation."""
|
||||||
return f"""Analyze this protocol: "{protocol_descriptor}"
|
return f"""Analyze this protocol: "{protocol_descriptor}"
|
||||||
@@ -153,6 +163,8 @@ Example usage:
|
|||||||
default='../bicorder.json',
|
default='../bicorder.json',
|
||||||
help='Path to bicorder.json (default: ../bicorder.json)')
|
help='Path to bicorder.json (default: ../bicorder.json)')
|
||||||
parser.add_argument('-m', '--model', help='LLM model to use')
|
parser.add_argument('-m', '--model', help='LLM model to use')
|
||||||
|
parser.add_argument('--resume', action='store_true',
|
||||||
|
help='Skip gradients that already have values in the CSV')
|
||||||
parser.add_argument('--dry-run', action='store_true',
|
parser.add_argument('--dry-run', action='store_true',
|
||||||
help='Show prompts without calling LLM or updating CSV')
|
help='Show prompts without calling LLM or updating CSV')
|
||||||
|
|
||||||
@@ -177,6 +189,15 @@ Example usage:
|
|||||||
bicorder_data = load_bicorder_config(args.bicorder)
|
bicorder_data = load_bicorder_config(args.bicorder)
|
||||||
gradients = extract_gradients(bicorder_data)
|
gradients = extract_gradients(bicorder_data)
|
||||||
|
|
||||||
|
# Load existing values for this row (for resume mode)
|
||||||
|
existing_values = get_row_values(args.csv_path, args.row_number) or {}
|
||||||
|
|
||||||
|
# Count existing values for reporting
|
||||||
|
if args.resume:
|
||||||
|
already_filled = sum(1 for g in gradients if existing_values.get(g['column_name'], '').strip())
|
||||||
|
if already_filled:
|
||||||
|
print(f"Resume: {already_filled}/{len(gradients)} gradients already have values")
|
||||||
|
|
||||||
if args.dry_run:
|
if args.dry_run:
|
||||||
print(f"DRY RUN: Row {args.row_number}, {len(gradients)} gradients")
|
print(f"DRY RUN: Row {args.row_number}, {len(gradients)} gradients")
|
||||||
print(f"Protocol: {protocol['descriptor']}\n")
|
print(f"Protocol: {protocol['descriptor']}\n")
|
||||||
@@ -188,6 +209,11 @@ Example usage:
|
|||||||
for i, gradient in enumerate(gradients, 1):
|
for i, gradient in enumerate(gradients, 1):
|
||||||
gradient_short = gradient['column_name'].replace('_', ' ')
|
gradient_short = gradient['column_name'].replace('_', ' ')
|
||||||
|
|
||||||
|
# In resume mode, skip gradients that already have a value
|
||||||
|
if args.resume and existing_values.get(gradient['column_name'], '').strip():
|
||||||
|
print(f"[{i}/{len(gradients)}] {gradient_short}: SKIP (already has value)")
|
||||||
|
continue
|
||||||
|
|
||||||
if not args.dry_run:
|
if not args.dry_run:
|
||||||
print(f"[{i}/{len(gradients)}] Querying: {gradient_short}...", flush=True)
|
print(f"[{i}/{len(gradients)}] Querying: {gradient_short}...", flush=True)
|
||||||
|
|
||||||
|
|||||||
Reference in new issue
Block a user