Team Ai
Datasetpublic

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.

sourceHugging Faceupdated 16d agoView on Hugging Face
18likes316downloads
classify-dataset.py629 linesDownload Raw Back to root
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