diff --git a/studio/backend/models/datasets.py b/studio/backend/models/datasets.py
index 73c75dc650..d891d099b6 100644
--- a/studio/backend/models/datasets.py
+++ b/studio/backend/models/datasets.py
@@ -44,6 +44,20 @@ class CheckFormatResponse(BaseModel):
warning: Optional[str] = None
+class AiAssistMappingRequest(BaseModel):
+ """Request for LLM-assisted column classification (user-triggered)."""
+ columns: List[str]
+ samples: List[Dict[str, Any]] # Preview rows already loaded in the dialog
+ dataset_name: Optional[str] = None # For LLM context
+
+
+class AiAssistMappingResponse(BaseModel):
+ """Response from LLM-assisted column classification."""
+ success: bool
+ suggested_mapping: Optional[Dict[str, str]] = None
+ warning: Optional[str] = None
+
+
class UploadDatasetResponse(BaseModel):
"""Response with stored dataset path for training."""
filename: str = Field(..., description="Original filename")
diff --git a/studio/backend/routes/datasets.py b/studio/backend/routes/datasets.py
index 62212754d4..318eef561c 100644
--- a/studio/backend/routes/datasets.py
+++ b/studio/backend/routes/datasets.py
@@ -36,6 +36,8 @@ if not logger.handlers:
from models.datasets import (
+ AiAssistMappingRequest,
+ AiAssistMappingResponse,
CheckFormatRequest,
CheckFormatResponse,
LocalDatasetItem,
@@ -479,3 +481,59 @@ def check_format(
status_code=500,
detail=f"Failed to check dataset format: {str(e)}"
)
+
+
+@router.post("/ai-assist-mapping", response_model=AiAssistMappingResponse)
+def ai_assist_mapping(
+ request: AiAssistMappingRequest,
+ current_subject: str = Depends(get_current_subject),
+):
+ """
+ Run LLM-assisted column classification on demand (user-triggered).
+
+ Receives the preview samples already loaded in the frontend dialog,
+ so no re-loading of the dataset is needed. The helper LLM is loaded
+ ephemerally: load → classify → unload.
+ """
+ try:
+ from utils.datasets.llm_assist import llm_classify_columns, llm_generate_dataset_warning
+
+ # Truncate sample values for the LLM prompt
+ truncated = [
+ {col: str(s.get(col, ""))[:200] for col in request.columns}
+ for s in request.samples[:5]
+ ]
+
+ mapping = llm_classify_columns(
+ column_names=request.columns,
+ samples=truncated,
+ )
+ if mapping:
+ # Keep only conversation roles, not metadata
+ conversation_mapping = {
+ col: role for col, role in mapping.items()
+ if role in ("user", "assistant", "system")
+ }
+ return AiAssistMappingResponse(
+ success=True,
+ suggested_mapping=conversation_mapping,
+ )
+
+ # LLM classification failed — generate a helpful warning
+ warning = llm_generate_dataset_warning(
+ issues=[f"Could not determine column roles from columns: {request.columns}"],
+ dataset_name=request.dataset_name,
+ modality="text",
+ column_names=request.columns,
+ )
+ return AiAssistMappingResponse(
+ success=False,
+ warning=warning or "AI could not determine column roles. Please assign them manually.",
+ )
+
+ except Exception as e:
+ logger.error(f"AI assist mapping failed: {e}", exc_info=True)
+ raise HTTPException(
+ status_code=500,
+ detail=f"AI assist failed: {str(e)}"
+ )
diff --git a/studio/backend/utils/datasets/dataset_utils.py b/studio/backend/utils/datasets/dataset_utils.py
index 35c7b796ec..5de54a387d 100644
--- a/studio/backend/utils/datasets/dataset_utils.py
+++ b/studio/backend/utils/datasets/dataset_utils.py
@@ -130,7 +130,7 @@ def check_dataset_format(dataset, is_vlm: bool = False) -> dict:
**audio_fields,
}
- # LLM flow
+ # Text / LLM flow
detected = detect_dataset_format(dataset)
# If format is unknown, try heuristic detection
@@ -149,66 +149,7 @@ def check_dataset_format(dataset, is_vlm: bool = False) -> dict:
**audio_fields,
}
else:
- # Heuristic failed — try LLM-assisted column classification
- try:
- from .llm_assist import llm_classify_columns
- from itertools import islice
-
- sample_rows = []
- for s in islice(dataset, 5):
- row = {col: str(s[col])[:200] for col in s}
- sample_rows.append(row)
-
- llm_mapping = llm_classify_columns(
- column_names=columns,
- samples=sample_rows,
- )
- if llm_mapping:
- # Keep only conversation roles, not metadata
- conversation_mapping = {
- col: role for col, role in llm_mapping.items()
- if role in ("user", "assistant", "system")
- }
- return {
- "requires_manual_mapping": False,
- "detected_format": "llm_assisted",
- "columns": columns,
- "suggested_mapping": conversation_mapping,
- "detected_image_column": None,
- "detected_text_column": None,
- "is_image": False,
- "multimodal_columns": None,
- **audio_fields,
- }
- except Exception as e:
- import logging
- logging.getLogger(__name__).debug(f"LLM column classification skipped: {e}")
-
- # Both heuristics and LLM failed — generate a meaningful warning
- warning = None
- try:
- from .llm_assist import llm_generate_dataset_warning
- from itertools import islice
- sample_rows = []
- for s in islice(dataset, 3):
- row = {col: str(s[col])[:150] for col in s}
- sample_rows.append(row)
- warning = llm_generate_dataset_warning(
- issues=[
- f"Could not auto-detect column roles from columns: {columns}",
- f"Sample values: {sample_rows[0] if sample_rows else 'N/A'}",
- ],
- modality="text",
- column_names=columns,
- )
- except Exception:
- pass
- if not warning:
- warning = (
- f"Could not auto-detect column roles for columns: {columns}. "
- "Please assign roles (user, assistant, etc.) manually."
- )
-
+ # Heuristic failed — user must map manually (or use AI Assist)
return {
"requires_manual_mapping": True,
"detected_format": "unknown",
@@ -218,7 +159,10 @@ def check_dataset_format(dataset, is_vlm: bool = False) -> dict:
"detected_text_column": None,
"is_image": False,
"multimodal_columns": None,
- "warning": warning,
+ "warning": (
+ f"Could not auto-detect column roles for columns: {columns}. "
+ "Please assign roles manually, or use AI Assist."
+ ),
**audio_fields,
}
diff --git a/studio/frontend/src/features/studio/sections/dataset-preview-dialog-mapping.tsx b/studio/frontend/src/features/studio/sections/dataset-preview-dialog-mapping.tsx
index 10a49be7fe..3eb8a1638f 100644
--- a/studio/frontend/src/features/studio/sections/dataset-preview-dialog-mapping.tsx
+++ b/studio/frontend/src/features/studio/sections/dataset-preview-dialog-mapping.tsx
@@ -14,6 +14,7 @@ import type { CheckFormatResponse } from "@/features/training/types/datasets";
import { cn } from "@/lib/utils";
import { AlertCircleIcon, CheckmarkCircle02Icon } from "@hugeicons/core-free-icons";
import { HugeiconsIcon } from "@hugeicons/react";
+import { Loader2, Sparkles } from "lucide-react";
const CHATML_ROLES = ["system", "user", "assistant"] as const;
const ALPACA_ROLES = ["instruction", "input", "output"] as const;
@@ -96,6 +97,9 @@ export function DatasetMappingCard({
isVlm = false,
isAudio = false,
format,
+ onAiAssist,
+ isAiLoading = false,
+ aiError,
}: {
mapping: Record
{aiError}
+ )} +