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
|
||||
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):
|
||||
"""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
|
||||
|
||||
|
||||
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."""
|
||||
cmd = ['python3', str(Path(__file__).parent / 'bicorder_query.py'), output_csv, str(row_num),
|
||||
'-b', bicorder_path]
|
||||
|
||||
if model:
|
||||
cmd.extend(['-m', model])
|
||||
if resume:
|
||||
cmd.extend(['--resume'])
|
||||
|
||||
print(f"Starting gradient queries...")
|
||||
|
||||
@@ -64,14 +69,28 @@ def query_gradients(output_csv, row_num, bicorder_path, model=None):
|
||||
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."""
|
||||
print(f"\n{'='*60}")
|
||||
print(f"Row {row_num}/{total_rows}")
|
||||
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)
|
||||
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")
|
||||
return False
|
||||
|
||||
@@ -112,7 +131,7 @@ Example usage:
|
||||
parser.add_argument('--end', type=int,
|
||||
help='End row number (1-indexed, default: all rows)')
|
||||
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()
|
||||
|
||||
@@ -156,7 +175,7 @@ Example usage:
|
||||
|
||||
for row_num in range(args.start, end_row + 1):
|
||||
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
|
||||
else:
|
||||
fail_count += 1
|
||||
|
||||
@@ -54,6 +54,16 @@ def get_protocol_by_row(csv_path, row_number):
|
||||
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):
|
||||
"""Generate a prompt for a single gradient evaluation."""
|
||||
return f"""Analyze this protocol: "{protocol_descriptor}"
|
||||
@@ -153,6 +163,8 @@ Example usage:
|
||||
default='../bicorder.json',
|
||||
help='Path to bicorder.json (default: ../bicorder.json)')
|
||||
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',
|
||||
help='Show prompts without calling LLM or updating CSV')
|
||||
|
||||
@@ -177,6 +189,15 @@ Example usage:
|
||||
bicorder_data = load_bicorder_config(args.bicorder)
|
||||
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:
|
||||
print(f"DRY RUN: Row {args.row_number}, {len(gradients)} gradients")
|
||||
print(f"Protocol: {protocol['descriptor']}\n")
|
||||
@@ -188,6 +209,11 @@ Example usage:
|
||||
for i, gradient in enumerate(gradients, 1):
|
||||
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:
|
||||
print(f"[{i}/{len(gradients)}] Querying: {gradient_short}...", flush=True)
|
||||
|
||||
|
||||
Reference in new issue
Block a user