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-sglang.py490 linesDownload Raw Back to root
1#!/usr/bin/env python32# /// script3# requires-python = ">=3.10"4# dependencies = [5#     "sglang[all]",6#     "flashinfer-python",7#     "transformers",8#     "torch",9#     "datasets",10#     "huggingface-hub[hf_transfer]",11# ]12# 13# [[tool.uv.index]]14# name = "flashinfer"15# url = "https://flashinfer.ai/whl/cu121/torch2.4/"16# ///17 18"""19Classify text columns in Hugging Face datasets using SGLang with reasoning-aware models.20 21This script provides efficient GPU-based classification with optional reasoning support,22optimized for models like SmolLM3-3B that use <think> tokens for chain-of-thought.23 24Example:25    # Fast classification without reasoning26    uv run classify-dataset-sglang.py \\27        --input-dataset imdb \\28        --column text \\29        --labels "positive,negative" \\30        --output-dataset user/imdb-classified31 32    # Complex classification with reasoning33    uv run classify-dataset-sglang.py \\34        --input-dataset arxiv-papers \\35        --column abstract \\36        --labels "reasoning_systems,agents,multimodal,robotics,other" \\37        --output-dataset user/arxiv-classified \\38        --reasoning39 40HF Jobs example:41    hf jobs uv run --flavor l4x1 \\42        https://huggingface.co/datasets/uv-scripts/classification/raw/main/classify-dataset-sglang.py \\43        --input-dataset user/emails \\44        --column content \\45        --labels "spam,ham" \\46        --output-dataset user/emails-classified \\47        --reasoning48"""49 50import argparse51import logging52import os53import sys54from typing import List, Dict, Any, Optional, Tuple55import json56import re57 58import torch59from datasets import load_dataset, Dataset60from huggingface_hub import HfApi, get_token61import sglang as sgl62 63# Default model - SmolLM3 with reasoning capabilities64DEFAULT_MODEL = "HuggingFaceTB/SmolLM3-3B"65 66# Minimum text length for valid classification67MIN_TEXT_LENGTH = 368 69# Maximum text length (in characters) to avoid context overflow70MAX_TEXT_LENGTH = 400071 72logging.basicConfig(73    level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s"74)75logger = logging.getLogger(__name__)76 77 78def parse_args():79    parser = argparse.ArgumentParser(80        description="Classify text in HuggingFace datasets using SGLang with reasoning support",81        formatter_class=argparse.RawDescriptionHelpFormatter,82        epilog=__doc__,83    )84 85    # Required arguments86    parser.add_argument(87        "--input-dataset",88        type=str,89        required=True,90        help="Input dataset ID on Hugging Face Hub",91    )92    parser.add_argument(93        "--column", type=str, required=True, help="Name of the text column to classify"94    )95    parser.add_argument(96        "--labels",97        type=str,98        required=True,99        help="Comma-separated list of classification labels (e.g., 'positive,negative')",100    )101    parser.add_argument(102        "--output-dataset",103        type=str,104        required=True,105        help="Output dataset ID on Hugging Face Hub",106    )107 108    # Optional arguments109    parser.add_argument(110        "--model",111        type=str,112        default=DEFAULT_MODEL,113        help=f"Model to use for classification (default: {DEFAULT_MODEL})",114    )115    parser.add_argument(116        "--reasoning",117        action="store_true",118        help="Enable reasoning mode (allows model to think through complex cases)",119    )120    parser.add_argument(121        "--save-reasoning",122        action="store_true",123        help="Save reasoning traces to a separate column (requires --reasoning)",124    )125    parser.add_argument(126        "--max-samples",127        type=int,128        default=None,129        help="Maximum number of samples to process (for testing)",130    )131    parser.add_argument(132        "--hf-token",133        type=str,134        default=None,135        help="Hugging Face API token (default: auto-detect from HF_TOKEN env var or huggingface-cli login)",136    )137    parser.add_argument(138        "--split",139        type=str,140        default="train",141        help="Dataset split to process (default: train)",142    )143    parser.add_argument(144        "--temperature",145        type=float,146        default=0.1,147        help="Temperature for generation (default: 0.1)",148    )149    parser.add_argument(150        "--max-tokens",151        type=int,152        default=500,153        help="Maximum tokens to generate (default: 500 for reasoning, 50 for non-reasoning)",154    )155    parser.add_argument(156        "--batch-size",157        type=int,158        default=32,159        help="Batch size for processing (default: 32)",160    )161    parser.add_argument(162        "--grammar-backend",163        type=str,164        default="xgrammar",165        choices=["outlines", "xgrammar", "llguidance"],166        help="Grammar backend for structured outputs (default: xgrammar)",167    )168 169    return parser.parse_args()170 171 172def preprocess_text(text: str) -> str:173    """Preprocess text for classification."""174    if not text or not isinstance(text, str):175        return ""176 177    # Strip whitespace178    text = text.strip()179 180    # Truncate if too long181    if len(text) > MAX_TEXT_LENGTH:182        text = f"{text[:MAX_TEXT_LENGTH]}..."183 184    return text185 186 187def validate_text(text: str) -> bool:188    """Check if text is valid for classification."""189    return bool(text and len(text) >= MIN_TEXT_LENGTH)190 191 192def create_classification_prompt(text: str, labels: List[str], reasoning: bool) -> str:193    """Create a prompt for classification with optional reasoning mode."""194    if reasoning:195        system_prompt = "You are a helpful assistant that thinks step-by-step before answering."196    else:197        system_prompt = "You are a helpful assistant. /no_think"198    199    user_prompt = f"""Classify this text as one of: {', '.join(labels)}200 201Text: {text}202 203Classification:"""204    205    # Format as a conversation206    return f"<|system|>\n{system_prompt}\n<|user|>\n{user_prompt}\n<|assistant|>\n"207 208 209def create_ebnf_grammar(labels: List[str]) -> str:210    """Create an EBNF grammar that constrains output to one of the given labels."""211    # Escape any special characters in labels212    escaped_labels = [f'"{label}"' for label in labels]213    choices = ' | '.join(escaped_labels)214    return f"root ::= {choices}"215 216 217def parse_reasoning_output(output: str, label: str) -> Optional[str]:218    """Extract reasoning from output if present."""219    # Look for thinking tags220    if "<think>" in output and "</think>" in output:221        start = output.find("<think>")222        end = output.find("</think>") + len("</think>")223        reasoning = output[start:end]224        return reasoning225    return None226 227 228def classify_batch_with_sglang(229    engine: sgl.Engine,230    texts: List[str],231    labels: List[str],232    args: argparse.Namespace233) -> List[Dict[str, Any]]:234    """Classify texts using SGLang with optional reasoning."""235    236    # Prepare prompts237    prompts = []238    valid_indices = []239    240    for i, text in enumerate(texts):241        processed_text = preprocess_text(text)242        if validate_text(processed_text):243            prompt = create_classification_prompt(processed_text, labels, args.reasoning)244            prompts.append(prompt)245            valid_indices.append(i)246    247    if not prompts:248        return [{"label": None, "reasoning": None} for _ in texts]249    250    # Set max tokens based on reasoning mode251    max_tokens = args.max_tokens if args.reasoning else 50252    253    # Create EBNF grammar for label constraints254    ebnf_grammar = create_ebnf_grammar(labels)255    256    # Set up sampling parameters with EBNF constraint257    sampling_params = {258        "temperature": args.temperature,259        "max_new_tokens": max_tokens,260        "ebnf": ebnf_grammar,  # This ensures output is one of the valid labels261    }262    263    try:264        # Generate with structured output constraint265        outputs = engine.generate(prompts, sampling_params)266        267        # Process outputs268        results = [{"label": None, "reasoning": None} for _ in texts]269        270        for idx, output in enumerate(outputs):271            original_idx = valid_indices[idx]272            273            # The output text should be just the label due to EBNF constraint274            classification = output.text.strip().strip('"')  # Remove quotes if present275            276            # Extract reasoning if present and requested277            reasoning = None278            if args.reasoning and args.save_reasoning:279                # Get the full output including reasoning280                # Note: We need to check if SGLang provides access to full output with reasoning281                reasoning = parse_reasoning_output(output.text, classification)282            283            results[original_idx] = {284                "label": classification,285                "reasoning": reasoning286            }287        288        return results289        290    except Exception as e:291        logger.error(f"Error during batch classification: {e}")292        # Return None labels for all texts in case of error293        return [{"label": None, "reasoning": None} for _ in texts]294 295 296def main():297    args = parse_args()298 299    # Validate reasoning arguments300    if args.save_reasoning and not args.reasoning:301        logger.error("--save-reasoning requires --reasoning to be enabled")302        sys.exit(1)303 304    # Check authentication early305    logger.info("Checking authentication...")306    token = args.hf_token or (os.environ.get("HF_TOKEN") or get_token())307 308    if not token:309        logger.error("No authentication token found. Please either:")310        logger.error("1. Run 'huggingface-cli login'")311        logger.error("2. Set HF_TOKEN environment variable")312        logger.error("3. Pass --hf-token argument")313        sys.exit(1)314 315    # Validate token by checking who we are316    try:317        api = HfApi(token=token)318        user_info = api.whoami()319        logger.info(f"Authenticated as: {user_info['name']}")320    except Exception as e:321        logger.error(f"Authentication failed: {e}")322        logger.error("Please check your token is valid")323        sys.exit(1)324 325    # Check CUDA availability326    if not torch.cuda.is_available():327        logger.error("CUDA is not available. This script requires a GPU.")328        logger.error("Please run on a machine with GPU support or use HF Jobs.")329        sys.exit(1)330 331    logger.info(f"CUDA available. Using device: {torch.cuda.get_device_name(0)}")332 333    # Parse and validate labels334    labels = [label.strip() for label in args.labels.split(",")]335    if len(labels) < 2:336        logger.error("At least two labels are required for classification.")337        sys.exit(1)338    logger.info(f"Classification labels: {labels}")339 340    # Load dataset341    logger.info(f"Loading dataset: {args.input_dataset}")342    try:343        dataset = load_dataset(args.input_dataset, split=args.split)344 345        # Limit samples if specified346        if args.max_samples:347            dataset = dataset.select(range(min(args.max_samples, len(dataset))))348            logger.info(f"Limited dataset to {len(dataset)} samples")349 350        logger.info(f"Loaded {len(dataset)} samples from split '{args.split}'")351    except Exception as e:352        logger.error(f"Failed to load dataset: {e}")353        sys.exit(1)354 355    # Verify column exists356    if args.column not in dataset.column_names:357        logger.error(f"Column '{args.column}' not found in dataset.")358        logger.error(f"Available columns: {dataset.column_names}")359        sys.exit(1)360 361    # Extract texts362    texts = dataset[args.column]363 364    # Initialize SGLang Engine365    logger.info(f"Initializing SGLang Engine with model: {args.model}")366    logger.info(f"Reasoning mode: {'enabled' if args.reasoning else 'disabled (fast mode)'}")367    logger.info(f"Grammar backend: {args.grammar_backend}")368    369    try:370        # Determine reasoning parser based on model371        reasoning_parser = None372        if "smollm3" in args.model.lower() or "qwen" in args.model.lower():373            reasoning_parser = "qwen"  # Uses <think> tokens374        elif "deepseek-r1" in args.model.lower():375            reasoning_parser = "deepseek-r1"376        377        engine_kwargs = {378            "model_path": args.model,379            "trust_remote_code": True,380            "dtype": "auto",381            "grammar_backend": args.grammar_backend,382        }383        384        if reasoning_parser and args.reasoning:385            engine_kwargs["reasoning_parser"] = reasoning_parser386            logger.info(f"Using reasoning parser: {reasoning_parser}")387        388        engine = sgl.Engine(**engine_kwargs)389        logger.info("SGLang engine initialized successfully")390    except Exception as e:391        logger.error(f"Failed to initialize SGLang: {e}")392        sys.exit(1)393 394    # Process in batches395    logger.info(f"Starting classification with batch size {args.batch_size}...")396    all_results = []397    398    for i in range(0, len(texts), args.batch_size):399        batch_end = min(i + args.batch_size, len(texts))400        batch_texts = texts[i:batch_end]401        402        logger.info(f"Processing batch {i//args.batch_size + 1}/{(len(texts) + args.batch_size - 1)//args.batch_size}")403        404        batch_results = classify_batch_with_sglang(405            engine, batch_texts, labels, args406        )407        all_results.extend(batch_results)408 409    # Extract labels and reasoning410    all_labels = [r["label"] for r in all_results]411    all_reasoning = [r["reasoning"] for r in all_results] if args.save_reasoning else None412 413    # Add classifications to dataset414    dataset = dataset.add_column("classification", all_labels)415    416    # Add reasoning column if requested417    if args.save_reasoning and args.reasoning:418        dataset = dataset.add_column("reasoning", all_reasoning)419        logger.info("Added reasoning traces to dataset")420 421    # Calculate statistics422    valid_count = sum(1 for label in all_labels if label is not None)423    invalid_count = len(all_labels) - valid_count424    425    if invalid_count > 0:426        logger.warning(427            f"{invalid_count} texts were too short or invalid for classification"428        )429 430    # Show classification distribution431    label_counts = {label: all_labels.count(label) for label in labels}432    logger.info("Classification distribution:")433    for label, count in label_counts.items():434        percentage = count / len(all_labels) * 100 if all_labels else 0435        logger.info(f"  {label}: {count} ({percentage:.1f}%)")436    if invalid_count > 0:437        none_percentage = invalid_count / len(all_labels) * 100438        logger.info(f"  Invalid/Skipped: {invalid_count} ({none_percentage:.1f}%)")439 440    # Log success rate441    success_rate = (valid_count / len(all_labels) * 100) if all_labels else 0442    logger.info(f"Classification success rate: {success_rate:.1f}%")443 444    # Save to Hub445    logger.info(f"Pushing dataset to Hub: {args.output_dataset}")446    try:447        commit_msg = f"Add classifications using {args.model} with SGLang"448        if args.reasoning:449            commit_msg += " (reasoning mode)"450        451        dataset.push_to_hub(452            args.output_dataset,453            token=token,454            commit_message=commit_msg,455        )456        logger.info(457            f"Successfully pushed to: https://huggingface.co/datasets/{args.output_dataset}"458        )459    except Exception as e:460        logger.error(f"Failed to push to Hub: {e}")461        sys.exit(1)462 463    # Clean up464    engine.shutdown()465    logger.info("SGLang engine shutdown complete")466 467 468if __name__ == "__main__":469    if len(sys.argv) == 1:470        print("Example HF Jobs commands:")471        print("\n# Fast classification (no reasoning):")472        print("hf jobs uv run \\")473        print("  --flavor l4x1 \\")474        print("  https://huggingface.co/datasets/uv-scripts/classification/raw/main/classify-dataset-sglang.py \\")475        print("  --input-dataset stanfordnlp/imdb \\")476        print("  --column text \\")477        print("  --labels 'positive,negative' \\")478        print("  --output-dataset user/imdb-classified")479        print("\n# Complex classification with reasoning:")480        print("hf jobs uv run \\")481        print("  --flavor l4x1 \\")482        print("  https://huggingface.co/datasets/uv-scripts/classification/raw/main/classify-dataset-sglang.py \\")483        print("  --input-dataset arxiv-papers \\")484        print("  --column abstract \\")485        print("  --labels 'reasoning_systems,agents,multimodal,robotics,other' \\")486        print("  --output-dataset user/arxiv-classified \\")487        print("  --reasoning --save-reasoning")488        sys.exit(0)489 490    main()