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-gliner2.py473 linesDownload Raw Back to root
1# /// script2# requires-python = ">=3.11,<3.14"3# dependencies = [4#     "gliner2[local]==2.0.0",5#     "protobuf",6#     "sentencepiece",7#     "datasets>=4.0.0,<6",8#     "huggingface-hub",9# ]10#11# [tool.hf-jobs]12# flavor = "t4-small"13# timeout = "1h"14# secrets = ["HF_TOKEN"]15# ///16"""17Classify a text column of a Hub dataset with GLiNER2 — zero-shot, or with your fine-tuned model.18 19GLiNER2 is a small encoder (74M to 287M parameters) that reads the label names as part of its20input. That gives two ways to use this script:21 221. Zero-shot: pass the label names with --labels. No training and no LLM. A t4-small does about23   33 rows/s; cpu-basic works but manages about 1.4 rows/s, so keep CPU for a few hundred rows.24   For English text, try `--model fastino/GLiNER2.5-Decide`.252. Fine-tuned: pass --model with a repo produced by `train-gliner2.py`. The tasks and labels are26   read from the model repo, so no --labels flag is needed.27 28Zero-shot on HF Jobs:29 30    hf jobs uv run --flavor t4-small --timeout 1h --secrets HF_TOKEN \\31        https://huggingface.co/datasets/uv-scripts/classification/raw/main/classify-gliner2.py \\32        fancyzhx/ag_news username/ag-news-gliner2 \\33        --labels World Sports Business "Science and technology" --max-samples 100034 35With a fine-tuned model:36 37    hf jobs uv run --flavor t4-small --timeout 1h --secrets HF_TOKEN \\38        https://huggingface.co/datasets/uv-scripts/classification/raw/main/classify-gliner2.py \\39        biglam/blbooksgenre username/blbooks-genre-predictions \\40        --dataset-config title_genre_classifiction --text-column title \\41        --model username/gliner2-blbooks-genre42 43Output: the original columns, plus `predicted_<task>` (a label, or a list of labels for a44multi-label task) and `predicted_<task>_confidence` for every task. The output dataset is45PRIVATE unless you pass --public.46 47Pass `--timeout` to `hf jobs uv run` for a big dataset: CLIs older than 1.32 ignore the48[tool.hf-jobs] header above and stop the job after 30 minutes, before anything is pushed.49"""50 51import argparse52import json53import logging54import os55import shlex56import sys57import time58from collections import Counter59 60os.environ.setdefault("TQDM_DISABLE", "1")61 62import datasets63import torch64from datasets import Features, List, Value, load_dataset65from gliner2.classification import (66    ClassificationConfig,67    ClassificationSchema,68    Classifier,69)70from huggingface_hub import DatasetCard, HfApi, hf_hub_download, login71from huggingface_hub.utils import (72    EntryNotFoundError,73    RepositoryNotFoundError,74    disable_progress_bars,75)76 77 78def configure_logging() -> logging.Logger:79    """Keep Jobs logs readable: root at WARNING, only this script's logger at INFO."""80    logging.basicConfig(81        level=logging.WARNING,82        format="%(asctime)s | %(levelname)s | %(message)s",83        datefmt="%H:%M:%S",84    )85    for noisy in ("httpx", "urllib3", "filelock", "huggingface_hub"):86        logging.getLogger(noisy).setLevel(logging.WARNING)87    disable_progress_bars()88    if hasattr(datasets, "disable_progress_bars"):89        datasets.disable_progress_bars()90 91    script_logger = logging.getLogger("classify-gliner2")92    script_logger.setLevel(logging.INFO)93    return script_logger94 95 96logger = configure_logging()97 98SCRIPT_URL = (99    "https://huggingface.co/datasets/uv-scripts/classification/raw/main/classify-gliner2.py"100)101DEFAULT_MODEL = "fastino/gliner2.5-multi-v1"102 103# Written into the model repo by train-gliner2.py: the tasks and labels the model was trained on.104SCHEMA_FILENAME = "classification_schema.json"105 106# GLiNER2 puts label names into the model prompt verbatim and rejects these strings.107FORBIDDEN_IN_LABELS = ("(", ")", "[P]", "[L]", "[C]", "[E]", "[R]", "[DESCRIPTION]", "[EXAMPLE]", "[OUTPUT]")108 109 110def check_labels(labels: list) -> None:111    for label in labels:112        for token in FORBIDDEN_IN_LABELS:113            if token in label:114                sys.exit(115                    f"Label {label!r} contains {token!r}, which GLiNER2 does not allow in a label "116                    "name. Rephrase it, for example with a dash instead of brackets."117                )118    if len(set(labels)) != len(labels):119        sys.exit(f"--labels contains a duplicate: {labels}")120    if len(labels) < 2:121        sys.exit("Pass at least two --labels.")122 123 124def exit_model_not_found(model_id: str) -> None:125    sys.exit(126        f"Cannot read the model '{model_id}'. Check the repo ID. If the repo is private or gated, "127        "make sure HF_TOKEN has access to it."128    )129 130 131def check_model_access(api: HfApi, model_id: str) -> None:132    """Stop with a clear message, before loading any data, if the model repo cannot be read."""133    if os.path.isdir(model_id):134        return135    try:136        api.model_info(model_id)137    except RepositoryNotFoundError:138        exit_model_not_found(model_id)139 140 141def load_trained_tasks(model_id: str):142    """Read the tasks that train-gliner2.py recorded in the model repo, or return None."""143    local_file = os.path.join(model_id, SCHEMA_FILENAME)144    if os.path.isfile(local_file):145        path = local_file146    elif os.path.isdir(model_id):147        return None148    else:149        try:150            path = hf_hub_download(model_id, SCHEMA_FILENAME)151        except EntryNotFoundError:152            return None153        except RepositoryNotFoundError:154            # Also raised for a gated repo the token has not been granted.155            exit_model_not_found(model_id)156    with open(path) as handle:157        return json.load(handle)["tasks"]158 159 160def resolve_tasks(args) -> list:161    """Decide which tasks to run: --labels wins, otherwise the model repo's recorded tasks."""162    if args.labels:163        check_labels(args.labels)164        return [{"name": args.task_name, "labels": args.labels, "multi_label": args.multi_label}]165 166    tasks = load_trained_tasks(args.model)167    if tasks is None:168        sys.exit(169            f"No --labels given, and '{args.model}' has no {SCHEMA_FILENAME}. Pass the label "170            "names with --labels, or use a model trained with train-gliner2.py."171        )172    logger.info("Using the %d task(s) recorded in %s.", len(tasks), args.model)173    return tasks174 175 176def build_schema(tasks: list) -> ClassificationSchema:177    schema = ClassificationSchema()178    for task in tasks:179        if task["multi_label"]:180            schema.multi(task["name"], task["labels"])181        else:182            schema.single(task["name"], task["labels"])183    return schema184 185 186def label_counts_table(tasks: list, counts_by_task: dict, total: int) -> str:187    lines = ["| Task | Label | Rows | Share |", "|---|---|---|---|"]188    for task in tasks:189        for label, count in counts_by_task[task["name"]].most_common():190            lines.append(f"| `{task['name']}` | {label} | {count} | {count / total:.1%} |")191    return "\n".join(lines)192 193 194# The smallest Jobs flavor for each GPU, keyed by a fragment of the GPU's name. "L40" comes195# before "L4" because the first match wins.196GPU_NAME_TO_FLAVOR = {"T4": "t4-small", "A10G": "a10g-small", "L40": "l40sx1", "L4": "l4x1", "A100": "a100-large"}197 198 199def jobs_flavor() -> str:200    """Return the Jobs hardware flavor, or "" when it is not known.201 202    The docs say ACCELERATOR holds the flavor ("a10g-small"). On the t4-small and a10g-small203    jobs that tested this script it held a bare "gpu", which is not a valid --flavor. So use204    ACCELERATOR when it looks like a flavor, and otherwise name the smallest flavor that has205    this GPU. A larger flavor of the same GPU reproduces the same result.206    """207    hardware = os.environ.get("ACCELERATOR") or ""208    looks_like_flavor = "-" in hardware or any(character.isdigit() for character in hardware)209    if looks_like_flavor:210        return hardware211    if not torch.cuda.is_available():212        return ""213    gpu_name = torch.cuda.get_device_name(0)214    for fragment, flavor in GPU_NAME_TO_FLAVOR.items():215        if fragment in gpu_name:216            return flavor217    return ""218 219 220def build_reproduce_command(args) -> str:221    flavor = jobs_flavor() or "t4-small"222    parts = [223        f"hf jobs uv run --flavor {flavor} --timeout 1h --secrets HF_TOKEN \\",224        f"  {SCRIPT_URL} \\",225        f"  {shlex.quote(args.input_dataset)} {shlex.quote(args.output_dataset)}",226    ]227    flags = []228    if args.model != DEFAULT_MODEL:229        flags.append(f"--model {shlex.quote(args.model)}")230    if args.labels:231        quoted = " ".join(shlex.quote(label) for label in args.labels)232        flags.append(f"--labels {quoted}")233        if args.task_name != "label":234            flags.append(f"--task-name {shlex.quote(args.task_name)}")235        if args.multi_label:236            flags.append("--multi-label")237    if args.dataset_config:238        flags.append(f"--dataset-config {shlex.quote(args.dataset_config)}")239    if args.text_column != "text":240        flags.append(f"--text-column {shlex.quote(args.text_column)}")241    if args.split != "train":242        flags.append(f"--split {shlex.quote(args.split)}")243    if args.max_samples:244        flags.append(f"--max-samples {args.max_samples}")245    if args.max_text_chars != 2000:246        flags.append(f"--max-text-chars {args.max_text_chars}")247    if args.public:248        flags.append("--public")249    if flags:250        parts[-1] += " \\"251        parts.append("  " + " ".join(flags))252    return "\n".join(parts)253 254 255def build_card(args, tasks, counts_by_task, total, seconds, zero_shot: bool) -> str:256    """Dataset card with the canonical uv-scripts provenance stamp."""257    on_jobs = os.environ.get("JOB_ID") is not None258    hardware = jobs_flavor()259    if on_jobs:260        origin = "Produced on [Hugging Face Jobs](https://huggingface.co/docs/huggingface_hub/guides/jobs)"261        if hardware:262            origin += f" (`{hardware}`)"263    else:264        origin = "Generated"265 266    tags = ["uv-script", "gliner2", "text-classification"]267    if on_jobs:268        tags.append("hf-jobs")269    tag_lines = "\n".join(f"- {tag}" for tag in tags)270 271    if zero_shot:272        how = (273            "The model was used **zero-shot**: it was given only the label names and has never "274            "seen labelled examples of this task. Treat the labels as a first pass to review, "275            "not as ground truth."276        )277    else:278        how = (279            "The model was fine-tuned for these tasks. Its model card reports the held-out "280            "scores. They apply only where this data resembles the training data."281        )282 283    column_lines = []284    for task in tasks:285        kind = "list of labels" if task["multi_label"] else "one label"286        column_lines.append(f"- `predicted_{task['name']}`: {kind} from {task['labels']}")287        note = " (empty when no label was selected)" if task["multi_label"] else ""288        column_lines.append(f"- `predicted_{task['name']}_confidence`: model confidence in [0, 1]{note}")289    column_block = "\n".join(column_lines)290 291    return f"""---292tags:293{tag_lines}294---295 296# {args.output_dataset.split("/")[-1]}297 298[`{args.input_dataset}`](https://huggingface.co/datasets/{args.input_dataset}) (split `{args.split}`,299{total} rows) with the `{args.text_column}` column classified by300[`{args.model}`](https://huggingface.co/{args.model}), a [GLiNER2](https://github.com/fastino-ai/GLiNER2) model.301 302{how}303 304## Added columns305 306{column_block}307 308Texts were truncated to {args.max_text_chars} characters before classification.309The confidence is not calibrated. Check it against a labelled sample before you use it as a filter.310 311## Label distribution312 313{label_counts_table(tasks, counts_by_task, total)}314 315Classified {total} rows in {round(seconds)} seconds ({total / max(seconds, 1e-9):.0f} rows/s).316 317## Reproduction318 319{origin} with the [`classify-gliner2.py`]({SCRIPT_URL}) recipe from [uv-scripts](https://huggingface.co/uv-scripts). Run it yourself:320 321```bash322{build_reproduce_command(args)}323```324"""325 326 327def in_own_account(api: HfApi, repo_id: str) -> str:328    """A bare name ("my-model") means a repo in your own account: return "<username>/my-model"."""329    if "/" in repo_id:330        return repo_id331    return f"{api.whoami()['name']}/{repo_id}"332 333 334def main(args) -> None:335    token = args.hf_token or os.environ.get("HF_TOKEN")336    if not token:337        sys.exit("No HF token. Pass --hf-token or run with --secrets HF_TOKEN.")338    login(token=token)339 340    # push_to_hub(private=True) leaves an existing repo's visibility alone, so check before the work.341    api = HfApi(token=token)342    args.output_dataset = in_own_account(api, args.output_dataset)343    if not os.path.exists(args.model):344        args.model = in_own_account(api, args.model)345    output_exists = api.repo_exists(args.output_dataset, repo_type="dataset")346    if not args.public and output_exists and not api.repo_info(args.output_dataset, repo_type="dataset").private:347        sys.exit(348            f"{args.output_dataset} already exists and is public. Pass --public to push there "349            "anyway, or choose a new dataset name."350        )351 352    check_model_access(api, args.model)353    tasks = resolve_tasks(args)354    for task in tasks:355        logger.info("Task '%s': %s", task["name"], task["labels"])356 357    logger.info("Loading %s (split %s)", args.input_dataset, args.split)358    dataset = load_dataset(args.input_dataset, args.dataset_config, split=args.split)359    if args.text_column not in dataset.column_names:360        sys.exit(f"Text column '{args.text_column}' not found. Columns are: {dataset.column_names}.")361    for task in tasks:362        for column in (f"predicted_{task['name']}", f"predicted_{task['name']}_confidence"):363            if column in dataset.column_names:364                sys.exit(f"The dataset already has a '{column}' column. Pass a different --task-name.")365    if args.max_samples and len(dataset) > args.max_samples:366        dataset = dataset.select(range(args.max_samples))367    logger.info("Rows to classify: %d", len(dataset))368 369    device = "cuda" if torch.cuda.is_available() else "cpu"370    if device == "cpu":371        logger.warning("No GPU found; classifying on CPU. Expect about 1-2 rows per second on cpu-basic.")372    # from_pretrained(device=...) does not move the weights in gliner2 2.0.0; .to() does.373    classifier = Classifier.from_pretrained(args.model).to(device=device).eval()374    schema = build_schema(tasks)375    config = ClassificationConfig(batch_size=args.batch_size)376 377    counts_by_task = {task["name"]: Counter() for task in tasks}378    empty_texts = 0379 380    def classify_batch(batch: dict) -> dict:381        nonlocal empty_texts382        texts = []383        for value in batch[args.text_column]:384            text = "" if value is None else str(value)385            if not text.strip():386                empty_texts += 1387                # The model needs some input; a missing text gets a prediction we then blank out.388                text = "-"389            texts.append(text[: args.max_text_chars])390 391        results = classifier.batch_classify(texts, schema, config=config)392 393        new_columns = {}394        for task in tasks:395            name = task["name"]396            predictions = []397            confidences = []398            for value, result in zip(batch[args.text_column], results):399                if value is None or not str(value).strip():400                    predictions.append([] if task["multi_label"] else None)401                    confidences.append(None)402                    continue403                if task["multi_label"]:404                    labels = list(result.selected(name))405                    predictions.append(labels)406                    counts_by_task[name].update(labels or ["(none selected)"])407                else:408                    label = result.value(name)409                    predictions.append(label)410                    counts_by_task[name][label] += 1411                confidence = result.confidence(name)412                confidences.append(None if confidence is None else float(confidence))413            new_columns[f"predicted_{name}"] = predictions414            new_columns[f"predicted_{name}_confidence"] = confidences415        return new_columns416 417    # Declare the output types. Otherwise the first map batch sets them, and a batch where418    # every confidence is None (no label selected, or no text) types the column as null and419    # the next batch fails to write.420    output_features = Features(dataset.features)421    for task in tasks:422        name = task["name"]423        output_features[f"predicted_{name}"] = List(Value("string")) if task["multi_label"] else Value("string")424        output_features[f"predicted_{name}_confidence"] = Value("float64")425 426    started = time.time()427    # One map batch holds several model batches, so progress is logged at a useful rate.428    dataset = dataset.map(429        classify_batch,430        batched=True,431        batch_size=args.batch_size * 8,432        features=output_features,433        load_from_cache_file=False,434    )435    seconds = time.time() - started436    logger.info("Classified %d rows in %.0f seconds.", len(dataset), seconds)437    if empty_texts:438        logger.warning("%d rows had no text and were left unlabelled.", empty_texts)439    for task in tasks:440        logger.info("Task '%s' distribution: %s", task["name"], dict(counts_by_task[task["name"]].most_common(10)))441 442    dataset.push_to_hub(args.output_dataset, private=not args.public)443    card = build_card(args, tasks, counts_by_task, len(dataset), seconds, zero_shot=bool(args.labels))444    DatasetCard(card).push_to_hub(args.output_dataset, repo_type="dataset")445    logger.info("Pushed to https://huggingface.co/datasets/%s", args.output_dataset)446 447 448def parse_args():449    parser = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)450    parser.add_argument("input_dataset", help="Input dataset ID")451    parser.add_argument("output_dataset", help="Output dataset: a name for your own account (my-dataset) or a full ID (org/my-dataset)")452    parser.add_argument("--model", default=DEFAULT_MODEL, help=f"GLiNER2 model: a base checkpoint for zero-shot, or a train-gliner2.py output (default: {DEFAULT_MODEL})")453    parser.add_argument("--labels", nargs="+", help="Label names for zero-shot classification. Overrides the tasks recorded in the model repo.")454    parser.add_argument("--task-name", default="label", help="Name of the --labels task; sets the output column names (default: label)")455    parser.add_argument("--multi-label", action="store_true", help="With --labels: allow several labels, or none, per text")456    parser.add_argument("--dataset-config", help="Dataset config name")457    parser.add_argument("--text-column", default="text", help="Text column (default: text)")458    parser.add_argument("--split", default="train", help="Split to classify (default: train)")459    parser.add_argument("--max-samples", type=int, help="Classify only the first N rows")460    parser.add_argument("--max-text-chars", type=int, default=2000, help="Truncate texts to this many characters (default: 2000)")461    parser.add_argument("--batch-size", type=int, default=32, help="Model batch size (default: 32)")462    parser.add_argument("--public", action="store_true", help="Make the output dataset public (default: private)")463    parser.add_argument("--private", action="store_true", help="Accepted for older commands; private is now the default")464    parser.add_argument("--hf-token", help="HF token (or set HF_TOKEN)")465    args = parser.parse_args()466    if args.public and args.private:467        parser.error("Pass --public or --private, not both.")468    return args469 470 471if __name__ == "__main__":472    main(parse_args())473