Download classify.py from CGIAR/hierarchical-text-classification: direct link, hf CLI and curl.
- Browser
- Download file 6.96 kB
-
https://huggingface.co/spaces/CGIAR/hierarchical-text-classification/resolve/main/classify.py
- Command line
-
hf download hf://spaces/CGIAR/hierarchical-text-classification/classify.py
-
curl -L -o classify.py https://huggingface.co/spaces/CGIAR/hierarchical-text-classification/resolve/main/classify.py
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) |