← FlyteCONTENT HISTORY

Update to Flyte

Snapshot Sep 30, 2026 · 22:59 UTC · version 1.0.1

Collection source: not recorded for this historical snapshot.

WHAT CHANGED · RULE-BASED ANALYSIS

First saved snapshot

No earlier snapshot is available to establish a change.

Compare saved observations

Download comparison JSON
Full technical diff · 0 changed fields
Full snapshot data
{
  "description": "Handles ML workload patterns: model training, hyperparameter optimization, experiment tracking, model evaluation and selection, batch inference, real-time serving, and model monitoring. Use when the user wants to train models, run hyperparameter search, track experiments, evaluate models, do batch or real-time inference, or set up model monitoring. Trigger words: \"train\", \"training\", \"hyperparameter\", \"HPO\", \"experiment\", \"tracking\", \"evaluation\", \"inference\", \"batch inference\", \"model serving\", \"monitoring\", \"GPU\", \"PyTorch\", \"TensorFlow\", \"scikit-learn\", \"HuggingFace\", \"model\".",
  "included_files": [],
  "name": "flyte-sdk-ml",
  "skill_md_contents": "---\nname: flyte-sdk-ml\ndescription: 'Handles ML workload patterns: model training, hyperparameter optimization, experiment tracking, model evaluation and selection, batch inference, real-time serving, and model monitoring. Use when the user wants to train models, run hyperparameter search, track experiments, evaluate models, do batch or real-time inference, or set up model monitoring. Trigger words: \"train\", \"training\", \"hyperparameter\", \"HPO\", \"experiment\", \"tracking\", \"evaluation\", \"inference\", \"batch inference\", \"model serving\", \"monitoring\", \"GPU\", \"PyTorch\", \"TensorFlow\", \"scikit-learn\", \"HuggingFace\", \"model\".'\n---\n\n# Flyte 2 SDK ML Skill\n\nBuild ML training, HPO, evaluation, and inference pipelines with Flyte 2.\n\n## Grounding References\n\n| Resource | URL |\n|---|---|\n| Official docs | https://www.union.ai/docs/v2/flyte |\n| Docs index (LLMs) | https://www.union.ai/docs/v2/flyte/llms.txt |\n| SDK API reference | https://www.union.ai/docs/v2/union/api-reference/flyte-sdk/ |\n| CLI API reference | https://www.union.ai/docs/v2/union/api-reference/flyte-cli/ |\n| flyte-sdk source | https://github.com/flyteorg/flyte-sdk |\n| Example code | https://github.com/unionai/unionai-examples |\n| Flyte MCP tools | Available via the `flyte-cluster` and `flyte-docs` MCP servers |\n\n**Ground unfamiliar APIs in real examples.** When unsure of a current Flyte 2 API, or for a pattern not shown below, and the `flyte-docs` search tools are available, search them first — by exact symbol (`TaskEnvironment`, `flyte.io.File`, `map_task`), since matching is literal substring, not semantic — then adapt a real example rather than inventing one, and cite the file or section you pulled it from. (Flyte 2 is not `flytekit`; priors are often wrong.)\n\n## Model Training\n\n### PyTorch Training\n\n```python\nimport flyte\nimport flyte.io\n\nenv = flyte.TaskEnvironment(\n    name=\"training\",\n    image=flyte.Image.from_base(\"pytorch/pytorch:2.1-cuda12.1-cudnn8-devel\").with_pip_packages(\n        \"transformers\", \"datasets\", \"accelerate\",\n    ),\n)\n\n@env.task(\n    requests=flyte.Resources(\n        cpu=\"4\", memory=\"16Gi\", gpu=\"1\", gpu_model=\"nvidia-a10g\",\n    ),\n)\nasync def train(\n    train_data: flyte.io.File,\n    val_data: flyte.io.File,\n    hyperparams: dict,\n) -> flyte.io.File:\n    \"\"\"Train a model and save checkpoint.\"\"\"\n    import torch\n    from transformers import AutoModelForSequenceClassification, AutoTokenizer\n\n    # Load data\n    tokenizer = AutoTokenizer.from_pretrained(\"bert-base-uncased\")\n    model = AutoModelForSequenceClassification.from_pretrained(\n        \"bert-base-uncased\", num_labels=2\n    )\n\n    # Train\n    for epoch in range(hyperparams[\"epochs\"]):\n        # ... training loop ...\n        pass\n\n    # Save checkpoint\n    output_path = \"/tmp/model_checkpoint\"\n    model.save_pretrained(output_path)\n    tokenizer.save_pretrained(output_path)\n    return flyte.io.File(path=output_path)\n\n@env.task\nasync def main(\n    train_uri: str,\n    val_uri: str,\n    lr: float = 0.001,\n    batch_size: int = 32,\n    epochs: int = 3,\n) -> dict:\n    hyperparams = {\"lr\": lr, \"batch_size\": batch_size, \"epochs\": epochs}\n    checkpoint = await train(\n        train_data=flyte.io.File(path=train_uri),\n        val_data=flyte.io.File(path=val_uri),\n        hyperparams=hyperparams,\n    )\n    return {\"checkpoint\": checkpoint, \"hyperparams\": hyperparams}\n```\n\n### scikit-learn Training\n\n```python\nimport flyte\nimport flyte.io\n\nenv = flyte.TaskEnvironment(\n    name=\"sklearn-training\",\n    image=flyte.Image.from_debian_base(python_version=(3, 12)).with_pip_packages(\n        \"scikit-learn\", \"pandas\", \"polars\", \"joblib\",\n    ),\n)\n\n@env.task\nasync def train_sklearn(\n    train_data: flyte.io.DataFrame,\n    val_data: flyte.io.DataFrame,\n    model_type: str = \"random_forest\",\n) -> flyte.io.File:\n    \"\"\"Train a scikit-learn model.\"\"\"\n    from sklearn.ensemble import RandomForestClassifier, GradientBoostingClassifier\n    from sklearn.linear_model import LogisticRegression\n    import joblib\n\n    X_train = train_data.to_polars().drop(\"label\").to_numpy()\n    y_train = train_data.to_polars()[\"label\"].to_numpy()\n    X_val = val_data.to_polars().drop(\"label\").to_numpy()\n    y_val = val_data.to_polars()[\"label\"].to_numpy()\n\n    if model_type == \"random_forest\":\n        model = RandomForestClassifier(n_estimators=100)\n    elif model_type == \"gbm\":\n        model = GradientBoostingClassifier(n_estimators=100)\n    else:\n        model = LogisticRegression()\n\n    model.fit(X_train, y_train)\n    accuracy = model.score(X_val, y_val)\n\n    path = f\"/tmp/{model_type}_model.joblib\"\n    joblib.dump(model, path)\n    return flyte.io.File(path=path)\n```\n\n### HuggingFace Trainer\n\n```python\nimport flyte\nimport flyte.io\n\nenv = flyte.TaskEnvironment(\n    name=\"hf-training\",\n    image=flyte.Image.from_base(\"pytorch/pytorch:2.1-cuda12.1-cudnn8-devel\").with_pip_packages(\n        \"transformers\", \"datasets\", \"accelerate\", \"evaluate\",\n    ),\n)\n\n@env.task(\n    requests=flyte.Resources(\n        cpu=\"4\", memory=\"16Gi\", gpu=\"1\", gpu_model=\"nvidia-a10g\",\n    ),\n)\nasync def train_hf(\n    dataset_name: str,\n    model_name: str,\n    hyperparams: dict,\n) -> flyte.io.File:\n    \"\"\"Train with HuggingFace Trainer.\"\"\"\n    from datasets import load_dataset\n    from transformers import (\n        AutoModelForSequenceClassification,\n        AutoTokenizer,\n        Trainer,\n        TrainingArguments,\n    )\n\n    train_dataset = load_dataset(dataset_name, split=\"train\")\n    val_dataset = load_dataset(dataset_name, split=\"validation\")\n\n    tokenizer = AutoTokenizer.from_pretrained(model_name)\n    model = AutoModelForSequenceClassification.from_pretrained(\n        model_name, num_labels=2\n    )\n\n    def tokenize(examples):\n        return tokenizer(examples[\"text\"], truncation=True, padding=\"max_length\", max_length=512)\n\n    train_dataset = train_dataset.map(tokenize)\n    val_dataset = val_dataset.map(tokenize)\n\n    training_args = TrainingArguments(\n        output_dir=\"/tmp/training_output\",\n        learning_rate=hyperparams.get(\"lr\", 2e-5),\n        per_device_train_batch_size=hyperparams.get(\"batch_size\", 16),\n        num_train_epochs=hyperparams.get(\"epochs\", 3),\n        evaluation_strategy=\"epoch\",\n        save_strategy=\"epoch\",\n    )\n\n    trainer = Trainer(\n        model=model,\n        args=training_args,\n        train_dataset=train_dataset,\n        eval_dataset=val_dataset,\n    )\n\n    trainer.train()\n    trainer.save_model(\"/tmp/final_model\")\n    tokenizer.save_pretrained(\"/tmp/final_model\")\n\n    return flyte.io.File(path=\"/tmp/final_model\")\n```\n\n## Hyperparameter Optimization\n\n### Manual HPO with fan-out\n\n```python\nimport flyte\n\nenv = flyte.TaskEnvironment(\n    name=\"hpo\",\n    image=flyte.Image.from_base(\"pytorch/pytorch:2.1-cuda12.1-cudnn8-devel\").with_pip_packages(\n        \"transformers\", \"datasets\",\n    ),\n)\n\n@env.task(\n    requests=flyte.Resources(\n        cpu=\"4\", memory=\"16Gi\", gpu=\"1\", gpu_model=\"nvidia-a10g\",\n    ),\n)\nasync def train_trial(hyperparams: dict) -> dict:\n    \"\"\"Run a single hyperparameter trial.\"\"\"\n    # hyperparams = {\"model\": \"bert-base\", \"lr\": 2e-5, \"batch_size\": 16, \"epochs\": 3}\n    checkpoint = await train_hf(\n        dataset_name=\"glue/mnli\",\n        model_name=hyperparams[\"model\"],\n        hyperparams=hyperparams,\n    )\n    # Evaluate\n    metrics = await evaluate(checkpoint, \"glue/mnli\", split=\"validation\")\n    return {\n        \"hyperparams\": hyperparams,\n        \"accuracy\": metrics[\"accuracy\"],\n        \"checkpoint\": checkpoint,\n    }\n\n@env.task\nasync def hpo_search(\n    param_grid: list[dict],\n) -> dict:\n    \"\"\"Run hyperparameter search with parallel trials.\"\"\"\n    # Fan out all trials in parallel\n    results = await flyte.map(train_trial, param_grid)\n    best = max(results, key=lambda r: r[\"accuracy\"])\n    return best\n\n@env.task\nasync def main() -> dict:\n    param_grid = [\n        {\"model\": \"bert-base\", \"lr\": 1e-5, \"batch_size\": 16, \"epochs\": 3},\n        {\"model\": \"bert-base\", \"lr\": 2e-5, \"batch_size\": 16, \"epochs\": 3},\n        {\"model\": \"bert-base\", \"lr\": 5e-5, \"batch_size\": 16, \"epochs\": 3},\n        {\"model\": \"bert-base\", \"lr\": 2e-5, \"batch_size\": 32, \"epochs\": 3},\n    ]\n    return await hpo_search(param_grid)\n```\n\n### Grid search pattern\n\n```python\nfrom itertools import product\n\n@env.task\nasync def grid_search() -> dict:\n    \"\"\"Grid search over hyperparameter combinations.\"\"\"\n    lr_values = [1e-5, 2e-5, 5e-5]\n    batch_sizes = [16, 32]\n    epochs = [2, 3]\n\n    param_grid = [\n        {\"model\": \"bert-base\", \"lr\": lr, \"batch_size\": bs, \"epochs\": ep}\n        for lr, bs, ep in product(lr_values, batch_sizes, epochs)\n    ]\n\n    results = await flyte.map(train_trial, param_grid)\n    best = max(results, key=lambda r: r[\"accuracy\"])\n    return best\n```\n\n## Experiment Tracking\n\n### Manual experiment tracking\n\n```python\nimport json\nimport datetime\nimport flyte\nimport flyte.io\n\n@env.task\nasync def track_experiment(\n    experiment_name: str,\n    hyperparams: dict,\n    metrics: dict,\n    checkpoint: flyte.io.File,\n) -> flyte.io.File:\n    \"\"\"Track experiment results as a JSON file in remote storage.\"\"\"\n    record = {\n        \"experiment\": experiment_name,\n        \"timestamp\": datetime.datetime.now().isoformat(),\n        \"hyperparameters\": hyperparams,\n        \"metrics\": metrics,\n        \"checkpoint_uri\": checkpoint.path,\n    }\n    path = f\"/tmp/experiments/{experiment_name}_{datetime.datetime.now().strftime('%Y%m%d_%H%M%S')}.json\"\n    with open(path, \"w\") as f:\n        json.dump(record, f, indent=2)\n    return flyte.io.File(path=path)\n\n@env.task\nasync def compare_experiments(\n    experiment_names: list[str],\n) -> dict:\n    \"\"\"Compare multiple experiments.\"\"\"\n    reports = []\n    for name in experiment_names:\n        report = await load_experiment(name)\n        reports.append(report)\n\n    # Find best by metric\n    best = max(reports, key=lambda r: r[\"metrics\"].get(\"accuracy\", 0))\n    return {\"best_experiment\": best, \"all\": reports}\n```\n\n### Inference result tracking\n\n```python\n@env.task\nasync def track_inference(\n    model_uri: str,\n    test_data: flyte.io.File,\n    metrics: dict,\n) -> flyte.io.File:\n    \"\"\"Track inference results.\"\"\"\n    record = {\n        \"model_uri\": model_uri,\n        \"test_data\": test_data.path,\n        \"metrics\": metrics,\n        \"timestamp\": datetime.datetime.now().isoformat(),\n    }\n    path = f\"/tmp/inference/{model_uri.split('/')[-1]}_{datetime.datetime.now().strftime('%Y%m%d')}.json\"\n    with open(path, \"w\") as f:\n        json.dump(record, f, indent=2)\n    return flyte.io.File(path=path)\n```\n\n## Model Evaluation and Selection\n\n### Evaluation pipeline\n\n```python\nimport flyte\nimport flyte.io\n\nenv = flyte.TaskEnvironment(\n    name=\"evaluation\",\n    image=flyte.Image.from_debian_base(python_version=(3, 12)).with_pip_packages(\n        \"scikit-learn\", \"scipy\", \"pandas\", \"matplotlib\", \"seaborn\",\n    ),\n)\n\n@env.task\nasync def evaluate_model(\n    model_path: flyte.io.File,\n    test_data: flyte.io.DataFrame,\n) -> dict:\n    \"\"\"Evaluate a model and return metrics.\"\"\"\n    import joblib\n    from sklearn.metrics import (\n        accuracy_score, f1_score, precision_score, recall_score,\n        roc_auc_score, confusion_matrix, classification_report,\n    )\n\n    model = joblib.load(model_path.path)\n    X_test = test_data.to_polars().drop(\"label\").to_numpy()\n    y_test = test_data.to_polars()[\"label\"].to_numpy()\n\n    y_pred = model.predict(X_test)\n    y_prob = model.predict_proba(X_test)[:, 1] if hasattr(model, \"predict_proba\") else y_pred\n\n    return {\n        \"accuracy\": accuracy_score(y_test, y_pred),\n        \"f1\": f1_score(y_test, y_pred),\n        \"precision\": precision_score(y_test, y_pred),\n        \"recall\": recall_score(y_test, y_pred),\n        \"auc\": roc_auc_score(y_test, y_prob),\n        \"confusion_matrix\": confusion_matrix(y_test, y_pred).tolist(),\n        \"report\": classification_report(y_test, y_pred, output_dict=True),\n    }\n\n@env.task\nasync def select_best_model(\n    candidate_models: list[flyte.io.File],\n    test_data: flyte.io.DataFrame,\n) -> dict:\n    \"\"\"Evaluate all candidates and select the best.\"\"\"\n    evaluations = await flyte.map(\n        lambda m: evaluate_model(m, test_data),\n        candidate_models,\n    )\n    best = max(evaluations, key=lambda e: e[\"accuracy\"])\n    return {\"best_metrics\": best, \"all_evaluations\": evaluations}\n```\n\n### Model comparison report\n\n```python\n@env.task\nasync def generate_comparison_report(\n    evaluations: list[dict],\n    model_names: list[str],\n) -> flyte.io.File:\n    \"\"\"Generate a model comparison report.\"\"\"\n    import matplotlib.pyplot as plt\n    import pandas as pd\n\n    df = pd.DataFrame({\n        \"model\": model_names,\n        \"accuracy\": [e[\"accuracy\"] for e in evaluations],\n        \"f1\": [e[\"f1\"] for e in evaluations],\n        \"precision\": [e[\"precision\"] for e in evaluations],\n        \"recall\": [e[\"recall\"] for e in evaluations],\n        \"auc\": [e[\"auc\"] for e in evaluations],\n    })\n\n    fig, axes = plt.subplots(1, 3, figsize=(18, 5))\n    metrics = [\"accuracy\", \"f1\", \"precision\", \"recall\", \"auc\"]\n    for i, metric in enumerate(metrics[:3]):\n        axes[i].bar(df[\"model\"], df[metric])\n        axes[i].set_title(metric)\n        axes[i].tick_params(axis=\"x\", rotation=45)\n\n    path = \"/tmp/model_comparison.png\"\n    fig.savefig(path, bbox_inches=\"tight\")\n    return flyte.io.File(path=path)\n```\n\n## Batch Inference\n\n### Large-scale batch inference\n\n```python\nimport flyte\nimport flyte.io\n\nenv = flyte.TaskEnvironment(\n    name=\"batch-inference\",\n    image=flyte.Image.from_debian_base(python_version=(3, 12)).with_pip_packages(\n        \"torch\", \"transformers\", \"pandas\", \"polars\", \"boto3\",\n    ),\n)\n\n@env.task(\n    requests=flyte.Resources(\n        cpu=\"4\", memory=\"16Gi\", gpu=\"1\", gpu_model=\"nvidia-a10g\",\n    ),\n)\nasync def load_model(model_uri: str) -> object:\n    \"\"\"Load model into memory.\"\"\"\n    from transformers import AutoModelForSequenceClassification, AutoTokenizer\n    tokenizer = AutoTokenizer.from_pretrained(model_uri)\n    model = AutoModelForSequenceClassification.from_pretrained(model_uri)\n    model.eval()\n    return {\"model\": model, \"tokenizer\": tokenizer}\n\n@env.task(\n    requests=flyte.Resources(\n        cpu=\"2\", memory=\"8Gi\", gpu=\"1\", gpu_model=\"nvidia-a10g\",\n    ),\n)\nasync def batch_predict(\n    model_ctx: object,\n    data_file: flyte.io.File,\n    batch_size: int = 32,\n) -> flyte.io.File:\n    \"\"\"Run inference on a batch of data.\"\"\"\n    import torch\n    import polars as pl\n\n    model = model_ctx[\"model\"]\n    tokenizer = model_ctx[\"tokenizer\"]\n\n    df = pl.read_parquet(data_file.path)\n    texts = df[\"text\"].to_list()\n\n    all_preds = []\n    all_probs = []\n    for i in range(0, len(texts), batch_size):\n        batch = texts[i:i + batch_size]\n        inputs = tokenizer(batch, padding=True, truncation=True, return_tensors=\"pt\")\n        with torch.no_grad():\n            outputs = model(**inputs)\n        probs = torch.softmax(outputs.logits, dim=1)\n        preds = torch.argmax(probs, dim=1)\n        all_preds.extend(preds.tolist())\n        all_probs.extend(probs.tolist())\n\n    results = pl.DataFrame({\"prediction\": all_preds, \"probability\": all_probs})\n    path = f\"/tmp/predictions_{data_file.path.split('/')[-1]}\"\n    results.write_parquet(path)\n    return flyte.io.File(path=path)\n\n@env.task\nasync def batch_inference(\n    model_uri: str,\n    data_files: list[str],\n) -> list:\n    \"\"\"Run batch inference on multiple data files.\"\"\"\n    model_ctx = await load_model(model_uri)\n    # Fan out inference across files\n    results = await flyte.map(\n        lambda f: batch_predict(model_ctx, flyte.io.File(path=f)),\n        data_files,\n    )\n    return results\n```\n\n### GPU batch inference optimization\n\n```python\n@env.task\nasync def optimized_batch_inference(\n    model_uri: str,\n    data_files: list[str],\n) -> list:\n    \"\"\"Optimized batch inference with dynamic batching.\"\"\"\n    # Use dynamic batcher for better GPU utilization\n    # Combine small batches and shard large ones\n    ...\n```\n\n## Real-time Model Serving\n\n### FastAPI model serving (covered in flyte-sdk-app)\n\n```python\nfrom fastapi import FastAPI\nimport flyte\nfrom flyte.app.extras import FastAPIAppEnvironment\n\napp = FastAPI()\nmodel = None\n\n@app.on_event(\"startup\")\nasync def load_model():\n    global model\n    from transformers import AutoModelForSequenceClassification, AutoTokenizer\n    model = AutoModelForSequenceClassification.from_pretrained(\"model-checkpoint\")\n    model.tokenizer = AutoTokenizer.from_pretrained(\"model-checkpoint\")\n\n@app.get(\"/predict\")\nasync def predict(text: str) -> dict:\n    inputs = model.tokenizer(text, return_tensors=\"pt\", padding=True, truncation=True)\n    with torch.no_grad():\n        outputs = model(**inputs)\n    probs = torch.softmax(outputs.logits, dim=1)\n    return {\n        \"prediction\": int(torch.argmax(probs, dim=1)[0]),\n        \"confidence\": float(probs.max().item()),\n    }\n\nenv = FastAPIAppEnvironment(\n    name=\"model-serving\",\n    app=app,\n    image=flyte.Image.from_base(\"pytorch/pytorch:2.1-cuda12.1-cudnn8-devel\").with_pip_packages(\n        \"fastapi\", \"uvicorn\", \"torch\", \"transformers\",\n    ),\n    resources=flyte.Resources(cpu=\"4\", memory=\"16Gi\", gpu=\"1\", gpu_model=\"nvidia-a10g\"),\n)\n```\n\n## Model Monitoring\n\n### Drift detection\n\n```python\n@env.task(cache=\"auto\")\nasync def detect_drift(\n    baseline_data: flyte.io.DataFrame,\n    current_data: flyte.io.DataFrame,\n) -> dict:\n    \"\"\"Detect data drift between baseline and current distributions.\"\"\"\n    import scipy.stats as stats\n\n    drift_results = {}\n    baseline_df = baseline_data.to_polars()\n    current_df = current_data.to_polars()\n\n    for col in baseline_df.columns:\n        if baseline_df[col].dtype.is_float64():\n            # Kolmogorov-Smirnov test\n            stat, p_value = stats.ks_2samp(\n                baseline_df[col].to_list(),\n                current_df[col].to_list(),\n            )\n            drift_results[col] = {\n                \"statistic\": stat,\n                \"p_value\": p_value,\n                \"drift_detected\": p_value < 0.05,\n            }\n\n    return drift_results\n\n@env.task\nasync def monitor_model(\n    model_uri: str,\n    baseline_data: flyte.io.DataFrame,\n    current_data: flyte.io.DataFrame,\n    predictions: flyte.io.DataFrame,\n) -> dict:\n    \"\"\"Monitor model health: drift, performance, prediction distribution.\"\"\"\n    drift = await detect_drift(baseline_data, current_data)\n\n    # Prediction distribution analysis\n    pred_dist = predictions.to_polars()[\"prediction\"].value_counts().to_dict()\n\n    # Confidence distribution\n    conf_stats = {\n        \"mean\": float(predictions.to_polars()[\"probability\"].mean()),\n        \"std\": float(predictions.to_polars()[\"probability\"].std()),\n        \"min\": float(predictions.to_polars()[\"probability\"].min()),\n        \"max\": float(predictions.to_polars()[\"probability\"].max()),\n    }\n\n    return {\n        \"drift\": drift,\n        \"prediction_distribution\": pred_dist,\n        \"confidence_stats\": conf_stats,\n        \"alert\": any(d[\"drift_detected\"] for d in drift.values()),\n    }\n```\n\n### Prediction quality monitoring\n\n```python\n@env.task\nasync def monitor_prediction_quality(\n    predictions: flyte.io.DataFrame,\n    ground_truth: flyte.io.DataFrame,\n) -> dict:\n    \"\"\"Monitor prediction quality over time.\"\"\"\n    merged = predictions.to_polars().join(ground_truth.to_polars(), on=\"id\")\n    accuracy = (merged[\"prediction\"] == merged[\"label\"]).mean()\n\n    # Per-class performance\n    per_class = {}\n    for label in merged[\"label\"].unique():\n        mask = merged[\"label\"] == label\n        per_class[int(label)] = {\n            \"count\": int(mask.sum()),\n            \"accuracy\": int(merged[mask][\"prediction\"] == merged[mask][\"label\"]).mean(),\n        }\n\n    return {\"accuracy\": float(accuracy), \"per_class\": per_class}\n```\n\n## ML Resource Recommendations\n\n| ML Workload | CPU | Memory | GPU |\n|---|---|---|---|\n| scikit-learn (small data) | 2-4 | 4-8 Gi | none |\n| scikit-learn (large data) | 4-8 | 16-32 Gi | none |\n| PyTorch training (small model) | 4 | 16 Gi | 1x A10G |\n| PyTorch training (large model) | 8 | 32+ Gi | 4-8x A100 |\n| HuggingFace fine-tuning | 4-8 | 16-32 Gi | 1-4x A10G/A100 |\n| Batch inference (CPU) | 4-8 | 16-32 Gi | none |\n| Batch inference (GPU) | 4 | 16 Gi | 1-4x A10G/A100 |\n| LLM serving | 8-16 | 32-64 Gi | 1-8x A100/H100 |\n\n## ML Anti-Patterns\n\n1. **Don't train without experiment tracking** — always log hyperparams, metrics, and model artifacts.\n2. **Don't skip evaluation** — always evaluate on held-out test data with multiple metrics.\n3. **Don't over-provision GPUs** — start with 1x A10G for most fine-tuning, scale only when needed.\n4. **Don't do batch inference one-by-one** — use `flyte.map` for parallel file-level inference.\n5. **Don't use Union-only features** — avoid `ReusePolicy` and other Union-specific APIs.\n6. **Don't forget to set `cache=\"auto\"`** on evaluation tasks — same model + same data = same result.\n"
}

SHA-256 of public snapshot: b42000af7c17a75c5297224165d94b1991657db160252c2892dfd089a14ff090