Nelly-43's picture
Update classify.py
45bc64e verified
Raw History Blame Contribute Delete
6.96 kB
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)