uv-scripts/classification
Classification Scripts Text classification on HF Jobs: label a dataset with a model that needs no training, or train your own classifier from labelled examples. If you have seen Jev and other "System One" models: the models these scripts train are small, open versions of the same idea. They read a piece of data and return a label with a probability, and you can train one on your own labels. For example, this demo suggests task tags for any Hub dataset; its model was fine-tuned… See the full description on the dataset page: https://huggingface.co/datasets/uv-scripts/classification.
18316
1#!/usr/bin/env python32# /// script3# requires-python = ">=3.10"4# dependencies = [5# "vllm>=0.6.6",6# "transformers>=4.53.0",7# "torch",8# "datasets",9# "huggingface-hub[hf_transfer]",10# ]11# ///12 13"""14Classify text columns in Hugging Face datasets using vLLM with structured outputs.15 16This script provides efficient GPU-based classification with guaranteed valid outputs,17optimized for running on HF Jobs.18 19Example:20 uv run classify-dataset.py \\21 --input-dataset imdb \\22 --column text \\23 --labels "positive,negative" \\24 --output-dataset user/imdb-classified25 26HF Jobs example:27 hfjobs run --flavor a10 uv run classify-dataset.py \\28 --input-dataset user/emails \\29 --column content \\30 --labels "spam,ham" \\31 --output-dataset user/emails-classified \\32 --prompt-style reasoning33"""34 35import argparse36import logging37import os38import sys39from typing import List40 41import torch42from datasets import load_dataset43from huggingface_hub import HfApi, get_token44from transformers import AutoTokenizer45from vllm import LLM, SamplingParams46from vllm.sampling_params import GuidedDecodingParams47 48# Default model - SmolLM3 for good balance of speed and quality49DEFAULT_MODEL = "HuggingFaceTB/SmolLM3-3B"50 51def parse_label_descriptions(desc_string: str) -> dict:52 """Parse label descriptions from CLI format 'label1:desc1,label2:desc2'."""53 if not desc_string:54 return {}55 56 descriptions = {}57 # Split by comma, but be careful about commas in descriptions58 parts = desc_string.split(',')59 60 current_label = None61 current_desc_parts = []62 63 for part in parts:64 if ':' in part and not current_label:65 # New label:description pair66 label, desc = part.split(':', 1)67 current_label = label.strip()68 current_desc_parts = [desc.strip()]69 elif ':' in part and current_label:70 # Save previous label and start new one71 descriptions[current_label] = ','.join(current_desc_parts)72 label, desc = part.split(':', 1)73 current_label = label.strip()74 current_desc_parts = [desc.strip()]75 else:76 # Continuation of previous description (had comma in it)77 current_desc_parts.append(part.strip())78 79 # Don't forget the last one80 if current_label:81 descriptions[current_label] = ','.join(current_desc_parts)82 83 return descriptions84 85 86def create_messages(text: str, labels: List[str], label_descriptions: dict = None, enable_reasoning: bool = False) -> List[dict]:87 """Create messages for chat template with optional label descriptions."""88 89 # Build the classification prompt90 if label_descriptions:91 # Format with descriptions92 categories_text = "Categories:\n"93 for label in labels:94 desc = label_descriptions.get(label, "")95 if desc:96 categories_text += f"- {label}: {desc}\n"97 else:98 categories_text += f"- {label}\n"99 else:100 # Simple format without descriptions101 categories_text = f"Categories: {', '.join(labels)}"102 103 if enable_reasoning:104 # Reasoning mode: allow thinking and request JSON output105 user_content = f"""Classify this text into one of these categories:106 107{categories_text}108 109Text: {text}110 111Think through your classification step by step, then provide your final answer in this JSON format:112{{"label": "your_chosen_label"}}"""113 114 system_content = "You are a helpful classification assistant that thinks step by step."115 else:116 # Structured output mode: fast classification117 if label_descriptions:118 user_content = f"Classify this text into one of these categories:\n\n{categories_text}\nText: {text}\n\nCategory:"119 else:120 user_content = f"Classify this text as one of: {', '.join(labels)}\n\nText: {text}\n\nLabel:"121 122 system_content = "You are a helpful classification assistant. /no_think"123 124 return [125 {"role": "system", "content": system_content},126 {"role": "user", "content": user_content}127 ]128 129# Minimum text length for valid classification130MIN_TEXT_LENGTH = 3131 132# Maximum text length (in characters) to avoid context overflow133MAX_TEXT_LENGTH = 4000134 135 136def parse_reasoning_output(output: str, valid_labels: List[str]) -> tuple[str, str, bool]:137 """Parse reasoning output to extract label from JSON after </think> tag.138 139 Returns:140 tuple: (label or None, full reasoning text, parsing_success)141 """142 import json143 144 # Find the </think> tag145 think_end = output.find("</think>")146 147 if think_end != -1:148 # Extract everything after </think>149 json_part = output[think_end + len("</think>"):].strip()150 reasoning = output[:think_end + len("</think>")]151 else:152 # No think tags, look for JSON in the output153 # Try to find JSON by looking for {154 json_start = output.find("{")155 if json_start != -1:156 json_part = output[json_start:].strip()157 reasoning = output[:json_start].strip() if json_start > 0 else ""158 else:159 json_part = output160 reasoning = output161 162 # Try to parse JSON163 try:164 # Find the first complete JSON object165 if "{" in json_part:166 # Extract just the JSON object167 json_str = json_part[json_part.find("{"):]168 # Find the matching closing brace169 brace_count = 0170 end_pos = 0171 for i, char in enumerate(json_str):172 if char == "{":173 brace_count += 1174 elif char == "}":175 brace_count -= 1176 if brace_count == 0:177 end_pos = i + 1178 break179 180 if end_pos > 0:181 json_str = json_str[:end_pos]182 data = json.loads(json_str)183 label = data.get("label", "")184 185 # Validate label186 if label in valid_labels:187 return label, output, True188 else:189 logger.warning(f"Parsed label '{label}' not in valid labels: {valid_labels}")190 return None, output, False191 else:192 logger.warning("Could not find complete JSON object")193 return None, output, False194 else:195 logger.warning("No JSON found in output")196 return None, output, False197 198 except json.JSONDecodeError as e:199 logger.warning(f"JSON parsing error: {e}")200 return None, output, False201 except Exception as e:202 logger.warning(f"Unexpected error parsing output: {e}")203 return None, output, False204 205logging.basicConfig(206 level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s"207)208logger = logging.getLogger(__name__)209 210 211def parse_args():212 parser = argparse.ArgumentParser(213 description="Classify text in HuggingFace datasets using vLLM with structured outputs",214 formatter_class=argparse.RawDescriptionHelpFormatter,215 epilog=__doc__,216 )217 218 # Required arguments219 parser.add_argument(220 "--input-dataset",221 type=str,222 required=True,223 help="Input dataset ID on Hugging Face Hub",224 )225 parser.add_argument(226 "--column", type=str, required=True, help="Name of the text column to classify"227 )228 parser.add_argument(229 "--labels",230 type=str,231 required=True,232 help="Comma-separated list of classification labels (e.g., 'positive,negative')",233 )234 parser.add_argument(235 "--output-dataset",236 type=str,237 required=True,238 help="Output dataset ID on Hugging Face Hub",239 )240 241 # Optional arguments242 parser.add_argument(243 "--model",244 type=str,245 default=DEFAULT_MODEL,246 help=f"Model to use for classification (default: {DEFAULT_MODEL})",247 )248 # Removed --batch-size argument as vLLM handles batching internally249 parser.add_argument(250 "--label-descriptions",251 type=str,252 default=None,253 help="Descriptions for each label in format 'label1:description1,label2:description2'",254 )255 parser.add_argument(256 "--enable-reasoning",257 action="store_true",258 help="Enable reasoning mode with thinking traces (disables structured outputs)",259 )260 parser.add_argument(261 "--max-samples",262 type=int,263 default=None,264 help="Maximum number of samples to process (for testing)",265 )266 parser.add_argument(267 "--hf-token",268 type=str,269 default=None,270 help="Hugging Face API token (default: auto-detect from HF_TOKEN env var or huggingface-cli login)",271 )272 parser.add_argument(273 "--split",274 type=str,275 default="train",276 help="Dataset split to process (default: train)",277 )278 parser.add_argument(279 "--temperature",280 type=float,281 default=0.1,282 help="Temperature for generation (default: 0.1)",283 )284 parser.add_argument(285 "--max-tokens",286 type=int,287 default=100,288 help="Maximum tokens to generate (default: 100, automatically increased 20x for reasoning mode)",289 )290 parser.add_argument(291 "--guided-backend",292 type=str,293 default="outlines",294 help="Guided decoding backend (default: outlines)",295 )296 parser.add_argument(297 "--shuffle",298 action="store_true",299 help="Shuffle dataset before selecting samples (useful with --max-samples for random sampling)",300 )301 parser.add_argument(302 "--shuffle-seed",303 type=int,304 default=42,305 help="Random seed for shuffling (default: 42)",306 )307 308 return parser.parse_args()309 310 311def preprocess_text(text: str) -> str:312 """Preprocess text for classification."""313 if not text or not isinstance(text, str):314 return ""315 316 # Strip whitespace317 text = text.strip()318 319 # Truncate if too long320 if len(text) > MAX_TEXT_LENGTH:321 text = f"{text[:MAX_TEXT_LENGTH]}..."322 323 return text324 325 326def validate_text(text: str) -> bool:327 """Check if text is valid for classification."""328 return bool(text and len(text) >= MIN_TEXT_LENGTH)329 330 331def prepare_prompts(332 texts: List[str], labels: List[str], tokenizer: AutoTokenizer, 333 label_descriptions: dict = None, enable_reasoning: bool = False334) -> tuple[List[str], List[int]]:335 """Prepare prompts using chat template for classification, filtering invalid texts."""336 prompts = []337 valid_indices = []338 339 for i, text in enumerate(texts):340 processed_text = preprocess_text(text)341 if validate_text(processed_text):342 # Create messages for chat template343 messages = create_messages(processed_text, labels, label_descriptions, enable_reasoning)344 345 # Apply chat template346 prompt = tokenizer.apply_chat_template(347 messages,348 tokenize=False,349 add_generation_prompt=True350 )351 prompts.append(prompt)352 valid_indices.append(i)353 354 return prompts, valid_indices355 356 357def main():358 args = parse_args()359 360 # Check authentication early361 logger.info("Checking authentication...")362 token = args.hf_token or (os.environ.get("HF_TOKEN") or get_token())363 364 if not token:365 logger.error("No authentication token found. Please either:")366 logger.error("1. Run 'huggingface-cli login'")367 logger.error("2. Set HF_TOKEN environment variable")368 logger.error("3. Pass --hf-token argument")369 sys.exit(1)370 371 # Validate token by checking who we are372 try:373 api = HfApi(token=token)374 user_info = api.whoami()375 logger.info(f"Authenticated as: {user_info['name']}")376 except Exception as e:377 logger.error(f"Authentication failed: {e}")378 logger.error("Please check your token is valid")379 sys.exit(1)380 381 # Check CUDA availability382 if not torch.cuda.is_available():383 logger.error("CUDA is not available. This script requires a GPU.")384 logger.error("Please run on a machine with GPU support or use HF Jobs.")385 sys.exit(1)386 387 logger.info(f"CUDA available. Using device: {torch.cuda.get_device_name(0)}")388 389 # Parse and validate labels390 labels = [label.strip() for label in args.labels.split(",")]391 if len(labels) < 2:392 logger.error("At least two labels are required for classification.")393 sys.exit(1)394 logger.info(f"Classification labels: {labels}")395 396 # Parse label descriptions if provided397 label_descriptions = None398 if args.label_descriptions:399 label_descriptions = parse_label_descriptions(args.label_descriptions)400 logger.info("Label descriptions provided:")401 for label, desc in label_descriptions.items():402 logger.info(f" {label}: {desc}")403 404 # Load dataset405 logger.info(f"Loading dataset: {args.input_dataset}")406 try:407 dataset = load_dataset(args.input_dataset, split=args.split)408 logger.info(f"Loaded {len(dataset)} samples from split '{args.split}'")409 410 # Shuffle if requested411 if args.shuffle:412 logger.info(f"Shuffling dataset with seed {args.shuffle_seed}")413 dataset = dataset.shuffle(seed=args.shuffle_seed)414 415 # Limit samples if specified416 if args.max_samples:417 dataset = dataset.select(range(min(args.max_samples, len(dataset))))418 logger.info(f"Limited dataset to {len(dataset)} samples")419 if args.shuffle:420 logger.info("Note: Samples were randomly selected due to shuffling")421 except Exception as e:422 logger.error(f"Failed to load dataset: {e}")423 sys.exit(1)424 425 # Verify column exists426 if args.column not in dataset.column_names:427 logger.error(f"Column '{args.column}' not found in dataset.")428 logger.error(f"Available columns: {dataset.column_names}")429 sys.exit(1)430 431 # Extract texts432 texts = dataset[args.column]433 434 # Load tokenizer for chat template formatting435 logger.info(f"Loading tokenizer for {args.model}")436 try:437 tokenizer = AutoTokenizer.from_pretrained(args.model, trust_remote_code=True)438 except Exception as e:439 logger.error(f"Failed to load tokenizer: {e}")440 sys.exit(1)441 442 # Initialize vLLM443 logger.info(f"Initializing vLLM with model: {args.model}")444 logger.info(f"Using guided decoding backend: {args.guided_backend}")445 try:446 llm = LLM(447 model=args.model,448 trust_remote_code=True,449 dtype="auto",450 gpu_memory_utilization=0.95,451 guided_decoding_backend=args.guided_backend,452 )453 except Exception as e:454 logger.error(f"Failed to initialize vLLM: {e}")455 sys.exit(1)456 457 # Set up sampling parameters based on mode458 if args.enable_reasoning:459 # Reasoning mode: no guided decoding, much more tokens for thinking460 sampling_params = SamplingParams(461 temperature=args.temperature,462 max_tokens=args.max_tokens * 20, # 20x more tokens for extensive reasoning463 )464 logger.info("Using reasoning mode - model will generate thinking traces with JSON output")465 else:466 # Structured output mode: guided decoding467 guided_params = GuidedDecodingParams(choice=labels)468 sampling_params = SamplingParams(469 guided_decoding=guided_params,470 temperature=args.temperature,471 max_tokens=args.max_tokens,472 )473 logger.info("Using structured output with guided_choice - outputs guaranteed to be valid labels")474 475 # Prepare all prompts476 logger.info("Preparing prompts for classification...")477 all_prompts, valid_indices = prepare_prompts(texts, labels, tokenizer, label_descriptions, args.enable_reasoning)478 479 if not all_prompts:480 logger.error("No valid texts found for classification.")481 sys.exit(1)482 483 logger.info(f"Prepared {len(all_prompts)} valid prompts out of {len(texts)} texts")484 485 # Let vLLM handle batching internally486 logger.info("Starting classification (vLLM will handle batching internally)...")487 488 try:489 # Generate all classifications at once - vLLM handles batching490 outputs = llm.generate(all_prompts, sampling_params)491 492 # Process outputs based on mode493 if args.enable_reasoning:494 # Reasoning mode: parse JSON and extract reasoning495 all_classifications = [None] * len(texts)496 all_reasoning = [None] * len(texts)497 all_parsing_success = [False] * len(texts)498 499 for idx, output in enumerate(outputs):500 original_idx = valid_indices[idx]501 generated_text = output.outputs[0].text.strip()502 503 # Parse the reasoning output504 label, reasoning, success = parse_reasoning_output(generated_text, labels)505 506 all_classifications[original_idx] = label507 all_reasoning[original_idx] = reasoning508 all_parsing_success[original_idx] = success509 510 # Log first few examples511 if idx < 3:512 logger.info(f"\nExample {idx + 1} output:")513 logger.info(f"Raw output: {generated_text[:200]}...")514 logger.info(f"Parsed label: {label}")515 logger.info(f"Parsing success: {success}")516 517 # Count parsing statistics518 parsing_success_count = sum(1 for s in all_parsing_success if s)519 parsing_fail_count = sum(1 for s in all_parsing_success if s is not None and not s)520 logger.info(f"\nParsing statistics:")521 logger.info(f" Successful: {parsing_success_count}/{len(valid_indices)} ({parsing_success_count/len(valid_indices)*100:.1f}%)")522 logger.info(f" Failed: {parsing_fail_count}/{len(valid_indices)} ({parsing_fail_count/len(valid_indices)*100:.1f}%)")523 524 valid_texts = parsing_success_count525 else:526 # Structured output mode: direct classification527 all_classifications = [None] * len(texts)528 for idx, output in enumerate(outputs):529 original_idx = valid_indices[idx]530 generated_text = output.outputs[0].text.strip()531 all_classifications[original_idx] = generated_text532 533 valid_texts = len(valid_indices)534 535 # Count statistics536 total_texts = len(texts)537 538 except Exception as e:539 logger.error(f"Classification failed: {e}")540 sys.exit(1)541 542 # Add columns to dataset543 dataset = dataset.add_column("classification", all_classifications)544 545 if args.enable_reasoning:546 dataset = dataset.add_column("reasoning", all_reasoning)547 dataset = dataset.add_column("parsing_success", all_parsing_success)548 549 # Calculate statistics550 none_count = total_texts - valid_texts551 if none_count > 0:552 logger.warning(553 f"{none_count} texts were too short or invalid for classification"554 )555 556 # Show classification distribution557 label_counts = {label: all_classifications.count(label) for label in labels}558 559 # Count None values separately560 none_classifications = all_classifications.count(None)561 562 logger.info("Classification distribution:")563 for label, count in label_counts.items():564 percentage = count / total_texts * 100 if total_texts > 0 else 0565 logger.info(f" {label}: {count} ({percentage:.1f}%)")566 567 if none_classifications > 0:568 none_percentage = none_classifications / total_texts * 100569 if args.enable_reasoning:570 logger.info(f" Failed to parse: {none_classifications} ({none_percentage:.1f}%)")571 else:572 logger.info(f" Invalid/Skipped: {none_classifications} ({none_percentage:.1f}%)")573 574 # Log success rate575 success_rate = (valid_texts / total_texts * 100) if total_texts > 0 else 0576 logger.info(f"Classification success rate: {success_rate:.1f}%")577 578 # Save to Hub (token already validated at start)579 logger.info(f"Pushing dataset to Hub: {args.output_dataset}")580 try:581 dataset.push_to_hub(582 args.output_dataset,583 token=token,584 commit_message=f"Add classifications using {args.model} {'with reasoning' if args.enable_reasoning else 'with structured outputs'}",585 )586 logger.info(587 f"Successfully pushed to: https://huggingface.co/datasets/{args.output_dataset}"588 )589 except Exception as e:590 logger.error(f"Failed to push to Hub: {e}")591 sys.exit(1)592 593 594if __name__ == "__main__":595 if len(sys.argv) == 1:596 print("Example commands:")597 print("\n# Simple classification:")598 print("uv run classify-dataset.py \\")599 print(" --input-dataset stanfordnlp/imdb \\")600 print(" --column text \\")601 print(" --labels 'positive,negative' \\")602 print(" --output-dataset user/imdb-classified")603 print("\n# With label descriptions:")604 print("uv run classify-dataset.py \\")605 print(" --input-dataset user/support-tickets \\")606 print(" --column content \\")607 print(" --labels 'bug,feature,question' \\")608 print(" --label-descriptions 'bug:something is broken or not working,feature:request for new functionality,question:asking for help or clarification' \\")609 print(" --output-dataset user/tickets-classified")610 print("\n# With reasoning mode (thinking + JSON output):")611 print("uv run classify-dataset.py \\")612 print(" --input-dataset stanfordnlp/imdb \\")613 print(" --column text \\")614 print(" --labels 'positive,negative,neutral' \\")615 print(" --enable-reasoning \\")616 print(" --output-dataset user/imdb-reasoned")617 print("\n# HF Jobs example:")618 print("hf jobs uv run \\")619 print(" --flavor l4x1 \\")620 print(" --image vllm/vllm-openai:latest \\")621 print(" classify-dataset.py \\")622 print(" --input-dataset stanfordnlp/imdb \\")623 print(" --column text \\")624 print(" --labels 'positive,negative' \\")625 print(" --output-dataset user/imdb-classified")626 sys.exit(0)627 628 main()629 