import os import time import tqdm import pandas as pd from concurrent.futures import ThreadPoolExecutor, as_completed from datetime import timedelta from llm import create_client, probe_endpoint, process_one_query from utils import setup_logging, load_system_prompt, load_queries, load_checkpoint, save_checkpoint def run_inference(cfg, df_input, ckpt_df, all_rows, total, future_to_idx, start_time, logger): completed = 0 # print("before: \n", df_input.columns) with tqdm.tqdm(total=total, desc="Classifying") as pbar: for future in as_completed(future_to_idx): try: idx, row = future.result() all_rows.append(row) except Exception as e: logger.error(f"Task failed: {e}") completed += 1 pbar.update(1) if completed % cfg.checkpoint_interval == 0: elapsed = time.time() - start_time elapsed_hours = elapsed / 3600 rate = completed / elapsed if elapsed > 0 else 0 # est_cost = elapsed_hours * COST_PER_HOUR remaining_sec = (total - completed) / rate if rate > 0 else 0 logger.info( f"Progress: {completed}/{total} " f"({completed/total*100:.1f}%) | " f"Rate: {rate:.1f} q/sec | " f"Elapsed: {timedelta(seconds=int(elapsed))} | " # f"Cost so far: ${est_cost:.2f} | " f"ETA: {timedelta(seconds=int(remaining_sec))}" ) batch_df = pd.concat([df_input, pd.DataFrame(all_rows)], axis=1) if ckpt_df is not None: combined = pd.concat([ckpt_df, batch_df], ignore_index=True) else: combined = batch_df save_checkpoint(combined, cfg) # print("after: \n", all_rows) # --- Final output --- final_df = pd.concat([df_input, pd.DataFrame(all_rows).drop(columns=[cfg.data[cfg.level].query_column])], axis=1) if ckpt_df is not None: final_df = pd.concat([ckpt_df, final_df], ignore_index=True) return final_df def classification_summary(cfg, start_time, final_df, logger): # --- Summary --- elapsed = time.time() - start_time elapsed_hours = elapsed / 3600 total_classified = len(final_df) valid = final_df["validation_ok"].sum() invalid = total_classified - valid # est_cost = elapsed_hours * COST_PER_HOUR print("valid: ", valid) print(final_df) logger.info("=" * 60) logger.info("CLASSIFICATION COMPLETE") logger.info(f" Total queries: {total_classified:,}") logger.info(f" Valid classifications: {valid:,} ({valid/total_classified*100:.1f}%)") logger.info(f" Validation failures: {invalid:,}") logger.info(f" Time elapsed: {timedelta(seconds=int(elapsed))}") # logger.info(f" Estimated cost: ${est_cost:.2f}") logger.info(f" Model: {cfg.model.model_name}") logger.info(f" Concurrency: {cfg.concurrency}") if valid > 0: conf_dist = final_df.loc[final_df["validation_ok"], "confidence"].value_counts() logger.info(" Confidence distribution:") for conf_level, count in conf_dist.items(): logger.info(f" {conf_level}: {count:,} ({count/valid*100:.1f}%)") # Commodity group match diagnostic # match_rate = final_df.loc[final_df["validation_ok"], "commodity_group_match"].mean() # mismatch_count = (~final_df.loc[final_df["validation_ok"], "commodity_group_match"]).sum() # logger.info(f" Commodity group match (LLM vs lookup): {match_rate*100:.1f}%") # if mismatch_count > 0: # logger.info(f" Commodity group mismatches: {mismatch_count:,} (check these for taxonomy drift)") logger.info("=" * 60) def process_queries(cfg, logger, start_time, client, system_prompt): data_config = cfg.data[cfg.level] df_input, queries = load_queries(data_config, logger) # df_input = df_input.iloc[:15] # queries = queries[:15] logger.info(f"Loaded {len(queries)} queries") if data_config.test_size != 0: queries = queries[:df_input.shape[0]] logger.info(f"TEST MODE: processing only {len(queries)} queries") # --- Check for checkpoint --- ckpt_df, num_done = load_checkpoint(cfg, logger) if num_done > 0 and num_done < len(queries): queries = queries[num_done:] logger.info(f"Remaining queries to classify: {len(queries)}") elif num_done >= len(queries): logger.info("All queries already classified. Nothing to do.") return else: logger.info("Starting fresh (no checkpoint found)") # --- Process queries concurrently --- all_rows = [] total = len(queries) tasks = [ (num_done + i, q, client, system_prompt, logger, cfg.model.model_name, data_config.query_column, data_config.classification_type, data_config.classification_fields) for i, q in enumerate(queries) ] logger.info(f"Starting classification: {total} queries, concurrency={cfg.concurrency}") # logger.info(f"Endpoint cost: ${COST_PER_HOUR}/hour while running") logger.info(f"Checkpointing every {cfg.checkpoint_interval} queries") with ThreadPoolExecutor(max_workers=cfg.concurrency) as executor: future_to_idx = { executor.submit(process_one_query, task): task[0] for task in tasks } final_df = run_inference(cfg, df_input, ckpt_df, all_rows, total, future_to_idx, start_time, logger) output_path = os.path.join(cfg.output_dir, cfg.output_file) final_df.to_csv(output_path, index=False) logger.info(f"Results saved to {output_path}") ckpt_path = os.path.join(cfg.output_dir, cfg.checkpoint_file) if os.path.exists(ckpt_path): os.remove(ckpt_path) logger.info("Checkpoint file removed (run complete)") classification_summary(cfg, start_time, final_df, logger) return output_path def classify_queries(cfg, logger): # os.makedirs(cfg.output_dir, exist_ok=True) # os.makedirs(cfg.exp_manager.log_dir, exist_ok=True) # logger = setup_logging(cfg.exp_manager.log_dir, cfg.exp_manager.log_file) # --- Load system prompt --- logger.info("Loading system prompt...") system_prompt = load_system_prompt(cfg.data[cfg.level].system_prompt_file) logger.info(f"System prompt loaded ({len(system_prompt)} chars)") # --- Create client --- logger.info(f"Connecting to: {cfg.model.endpoint_url}") logger.info(f"Model: {cfg.model.model_name}") client = create_client(cfg.model) # --- Probe mode --- if cfg.probe: probe_endpoint(client, system_prompt, logger, cfg.model) return start_time = time.time() return process_queries(cfg, logger, start_time, client, system_prompt)