#!/usr/bin/env python3 """Batch training data generator. Processes a folder of data files (.sas7bdat, .csv, .xpt) through an LLM to generate training examples for fine-tuning. Designed for use with public/open datasets (CDISC pilot, PhUSE, FDA submissions). Usage: # Process a folder of SAS files python scripts/batch_training.py ./open_data/cdisc_pilot/ # Process with a specific prompt python scripts/batch_training.py ./open_data/ --prompt "classify for demographic analysis" # Process with a specific model python scripts/batch_training.py ./open_data/ --model groq/llama-3.3-70b-versatile # Dry run — show what would be processed without calling the LLM python scripts/batch_training.py ./open_data/ --dry-run # Process files in groups (simulating multi-file upload) python scripts/batch_training.py ./open_data/ --group-by-prefix """ import argparse import json import logging import os import sys from pathlib import Path # Add project root to path sys.path.insert(0, str(Path(__file__).parent.parent)) from openai import OpenAI from app.core.sas_extractor import SUPPORTED_EXTENSIONS, extract_metadata from app.core.training_collector import save_training_example from app.models.mapping import ColumnClassification, FileJoin, MappingResult from app.models.metadata import DatasetMetadata from app.prompts.mapping_prompt import SYSTEM_PROMPT, build_user_prompt logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s") logger = logging.getLogger(__name__) DEFAULT_PROMPT = ( "Classify all columns for ingestion into a standard data + metadata template. " "Identify the sample/subject identifiers, clinical/demographic metadata, " "and measured data variables." ) def find_data_files(folder: Path) -> list[Path]: """Find all supported data files in a folder (recursive).""" files = [] for ext in SUPPORTED_EXTENSIONS: files.extend(folder.rglob(f"*{ext}")) return sorted(files) def group_by_prefix(files: list[Path]) -> list[list[Path]]: """Group files by common prefix (e.g., study_dm.sas7bdat, study_ae.sas7bdat). Heuristic: files in the same directory with a common prefix before an underscore or dot are grouped together. """ from collections import defaultdict groups = defaultdict(list) for f in files: # Group by parent directory groups[f.parent].append(f) return list(groups.values()) def classify_with_llm( all_metadata: list[DatasetMetadata], user_prompt: str, model: str, api_key: str | None = None, base_url: str | None = None, ) -> dict: """Send metadata to the LLM and get classifications back.""" from app.models.template import TemplateSchema template = TemplateSchema( template_name="aseesa_standard_v1", version="1.0", sample_id_field="Sample_ID", subject_id_field="Subject_title", common_metadata_fields=[], ) prompt = build_user_prompt(all_metadata, template, user_prompt) client = OpenAI( api_key=api_key or "not-set", base_url=base_url or "https://api.groq.com/openai/v1", ) response = client.chat.completions.create( model=model, max_tokens=8192, messages=[ {"role": "system", "content": SYSTEM_PROMPT}, {"role": "user", "content": prompt}, ], ) raw_text = response.choices[0].message.content # Parse JSON json_str = raw_text if "```json" in json_str: json_str = json_str.split("```json")[1].split("```")[0] elif "```" in json_str: json_str = json_str.split("```")[1].split("```")[0] return json.loads(json_str.strip()) def process_group( files: list[Path], user_prompt: str, model: str, api_key: str | None, base_url: str | None = None, dry_run: bool = False, ) -> bool: """Process a group of files and save as a training example.""" group_name = ", ".join(f.name for f in files) logger.info("Processing: %s", group_name) # Extract metadata all_metadata = [] for f in files: try: meta = extract_metadata(f) meta.filename = f.name # Use original name all_metadata.append(meta) except Exception as e: logger.warning(" Skipping %s: %s", f.name, e) continue if not all_metadata: logger.warning(" No valid files in group, skipping") return False # Show what we found total_cols = sum(m.num_columns for m in all_metadata) logger.info(" Found %d file(s), %d total columns", len(all_metadata), total_cols) if dry_run: for m in all_metadata: logger.info(" %s: %d cols, %d rows", m.filename, m.num_columns, m.num_rows) for col in m.columns[:5]: logger.info(" - %s (%s) %s", col.name, col.dtype, col.label or "") if m.num_columns > 5: logger.info(" ... and %d more", m.num_columns - 5) return True # Call LLM try: parsed = classify_with_llm(all_metadata, user_prompt, model, api_key, base_url) except Exception as e: logger.error(" LLM classification failed: %s", e) return False # Parse result try: columns = [ColumnClassification(**c) for c in parsed.get("columns", [])] joins = [FileJoin(**j) for j in parsed.get("joins", [])] primary_file = parsed.get("primary_file", all_metadata[0].filename) except Exception as e: logger.error(" Failed to parse LLM response: %s", e) return False # Validate basics sample_ids = [c for c in columns if c.classification == "sample_id"] if not sample_ids: logger.warning(" LLM didn't identify a sample_id, skipping") return False # Save training example session_id = f"batch_{files[0].stem}" save_training_example( session_id=session_id, all_metadata=all_metadata, columns=columns, joins=joins, primary_file=primary_file, ) logger.info( " Saved: %d columns classified, %d joins, primary=%s", len(columns), len(joins), primary_file, ) return True def main(): parser = argparse.ArgumentParser( description="Generate training data from a folder of data files" ) parser.add_argument("folder", type=Path, help="Folder containing data files") parser.add_argument("--prompt", default=DEFAULT_PROMPT, help="Analysis prompt for the LLM") parser.add_argument("--model", default=None, help="LLM model (default: from LLM_MODEL env)") parser.add_argument("--api-key", default=None, help="API key (default: from LLM_API_KEY env)") parser.add_argument("--base-url", default=None, help="Base URL (default: from LLM_BASE_URL env)") parser.add_argument("--group-by-prefix", action="store_true", help="Group files by directory (simulates multi-file upload)") parser.add_argument("--dry-run", action="store_true", help="Show what would be processed without calling LLM") args = parser.parse_args() if not args.folder.exists(): logger.error("Folder not found: %s", args.folder) sys.exit(1) model = args.model or os.environ.get("LLM_MODEL", "llama-3.3-70b-versatile") api_key = args.api_key or os.environ.get("LLM_API_KEY") base_url = args.base_url or os.environ.get("LLM_BASE_URL", "https://api.groq.com/openai/v1") if not api_key and not args.dry_run: logger.error("No API key. Set LLM_API_KEY env var or use --api-key") sys.exit(1) # Find files files = find_data_files(args.folder) if not files: logger.error("No supported files found in %s", args.folder) sys.exit(1) logger.info("Found %d data files in %s", len(files), args.folder) # Group files if args.group_by_prefix: groups = group_by_prefix(files) logger.info("Grouped into %d batches", len(groups)) else: # Each file is its own group groups = [[f] for f in files] # Process success = 0 failed = 0 for group in groups: ok = process_group(group, args.prompt, model, api_key, base_url, args.dry_run) if ok: success += 1 else: failed += 1 logger.info("Done: %d succeeded, %d failed", success, failed) if not args.dry_run: training_file = Path("./training_data/examples.jsonl") if training_file.exists(): count = sum(1 for _ in open(training_file)) logger.info("Total training examples: %d (in %s)", count, training_file) if __name__ == "__main__": main()