Spaces:
Sleeping
Sleeping
| #!/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() | |