From 0e3ac91e2a0e531c2ee199549bfe7990b882f1bd Mon Sep 17 00:00:00 2001 From: Roland Tannous Date: Wed, 11 Mar 2026 16:55:43 +0000 Subject: [PATCH] feat: target AI Assist mapping prompts for audio & embedding models --- studio/backend/models/datasets.py | 2 + studio/backend/routes/datasets.py | 2 + studio/backend/utils/datasets/llm_assist.py | 40 +++++++++++++++++++ .../sections/dataset-preview-dialog.tsx | 8 +++- .../src/features/training/api/datasets-api.ts | 6 +++ 5 files changed, 57 insertions(+), 1 deletion(-) diff --git a/studio/backend/models/datasets.py b/studio/backend/models/datasets.py index bc4f4d047f..f84c98fc0a 100644 --- a/studio/backend/models/datasets.py +++ b/studio/backend/models/datasets.py @@ -50,6 +50,8 @@ class AiAssistMappingRequest(BaseModel): samples: List[Dict[str, Any]] # Preview rows already loaded in the dialog dataset_name: Optional[str] = None # For LLM context hf_token: Optional[str] = None # For fetching dataset card + model_name: Optional[str] = None + model_type: Optional[str] = None class AiAssistMappingResponse(BaseModel): diff --git a/studio/backend/routes/datasets.py b/studio/backend/routes/datasets.py index 90df09811c..29f7cba005 100644 --- a/studio/backend/routes/datasets.py +++ b/studio/backend/routes/datasets.py @@ -506,6 +506,8 @@ def ai_assist_mapping( samples=truncated, dataset_name=request.dataset_name, hf_token=request.hf_token, + model_name=request.model_name, + model_type=request.model_type, ) if result and result.get("success"): diff --git a/studio/backend/utils/datasets/llm_assist.py b/studio/backend/utils/datasets/llm_assist.py index 2481a63ba5..f7016f5ff1 100644 --- a/studio/backend/utils/datasets/llm_assist.py +++ b/studio/backend/utils/datasets/llm_assist.py @@ -434,6 +434,9 @@ def _run_multi_pass_advisor( dataset_name: Optional[str] = None, dataset_card: Optional[str] = None, dataset_metadata: Optional[dict] = None, + model_name: Optional[str] = None, + model_type: Optional[str] = None, + hf_token: Optional[str] = None, ) -> Optional[dict[str, Any]]: """ Multi-pass LLM analysis: classify → convert → validate. @@ -480,6 +483,36 @@ def _run_multi_pass_advisor( ) card_excerpt = (dataset_card or "")[:1200] or "N/A" + # ── Target Model Hints ── + target_hints = "" + is_gemma_3n = False + if model_name: + try: + from utils.models.model_config import load_model_config + config = load_model_config(model_name, use_auth=True, token=hf_token) + archs = getattr(config, "architectures", []) + if archs and "Gemma3nForConditionalGeneration" in archs: + is_gemma_3n = True + except Exception: + is_gemma_3n = "gemma-3n" in model_name.lower() + + if model_type == "audio" and not is_gemma_3n: + target_hints = ( + "\n\nHINT: The user is training an AUDIO model. The dataset MUST contain " + "a column with audio files/paths. Ensure one such column is selected " + "as part of the input." + ) + elif model_type == "embeddings": + target_hints = ( + "\n\nHINT: The user is training an EMBEDDING model. These models typically " + "do not use standard conversational input/output formats but instead use " + "specific formats like:\n" + "- Pairs of texts for Semantic Textual Similarity (STS)\n" + "- Premise, hypothesis, and label for Natural Language Inference (NLI)\n" + "- Queries and positive/negative documents for information retrieval\n" + "Ensure the dataset format mapped reflects these specialized tasks." + ) + # ── Pass 1: Classify ── print("🤖 Pass 1: Classifying dataset...", flush=True) t1 = time.monotonic() @@ -495,6 +528,7 @@ def _run_multi_pass_advisor( "— they are things like summarization, question answering, translation, " "classification, etc. Those need conversion. You must respond with ONLY a " "valid JSON object. Do not write any explanation before or after the JSON." + f"{target_hints}" ), }, { @@ -568,6 +602,7 @@ def _run_multi_pass_advisor( '4. Metadata columns like "id", "index", "source", "url", "date" should be ' 'set to "skip".\n\n' "You must respond with ONLY a valid JSON object." + f"{target_hints}" ), }, { @@ -749,6 +784,8 @@ def llm_conversion_advisor( samples: list[dict], dataset_name: Optional[str] = None, hf_token: Optional[str] = None, + model_name: Optional[str] = None, + model_type: Optional[str] = None, ) -> Optional[dict[str, Any]]: """ Full conversion advisor: fetch HF card → multi-pass LLM analysis. @@ -773,6 +810,9 @@ def llm_conversion_advisor( dataset_name=dataset_name, dataset_card=dataset_card, dataset_metadata=dataset_metadata, + model_name=model_name, + model_type=model_type, + hf_token=hf_token, ) if result and result.get("success"): diff --git a/studio/frontend/src/features/studio/sections/dataset-preview-dialog.tsx b/studio/frontend/src/features/studio/sections/dataset-preview-dialog.tsx index 1cb9b4e741..c79cf77e90 100644 --- a/studio/frontend/src/features/studio/sections/dataset-preview-dialog.tsx +++ b/studio/frontend/src/features/studio/sections/dataset-preview-dialog.tsx @@ -69,6 +69,8 @@ export function DatasetPreviewDialog({ manualMapping, setManualMapping, datasetFormat, setDatasetAdvisorFields, datasetAdvisorNotification, datasetSystemPrompt, + selectedModel, + modelType, } = useTrainingConfigStore( useShallow((s) => ({ manualMapping: s.datasetManualMapping, @@ -77,6 +79,8 @@ export function DatasetPreviewDialog({ setDatasetAdvisorFields: s.setDatasetAdvisorFields, datasetAdvisorNotification: s.datasetAdvisorNotification, datasetSystemPrompt: s.datasetSystemPrompt, + selectedModel: s.selectedModel, + modelType: s.modelType, })), ); const { isStarting, startError, startTrainingRun } = useTrainingActions(); @@ -108,6 +112,8 @@ export function DatasetPreviewDialog({ samples: data.preview_samples, datasetName: datasetName, hfToken: hfToken, + modelName: selectedModel, + modelType: modelType, }); if (result.success && result.suggested_mapping) { @@ -135,7 +141,7 @@ export function DatasetPreviewDialog({ } finally { setIsAiLoading(false); } - }, [data, datasetFormat, datasetName, hfToken, setManualMapping, setDatasetAdvisorFields]); + }, [data, datasetFormat, datasetName, hfToken, setManualMapping, setDatasetAdvisorFields, selectedModel, modelType]); // When format changes, remap existing mapping roles to the new format's role names const prevFormatRef = useRef(datasetFormat); diff --git a/studio/frontend/src/features/training/api/datasets-api.ts b/studio/frontend/src/features/training/api/datasets-api.ts index 548e6ef2e9..3ff04189f7 100644 --- a/studio/frontend/src/features/training/api/datasets-api.ts +++ b/studio/frontend/src/features/training/api/datasets-api.ts @@ -69,6 +69,8 @@ type AiAssistMappingArgs = { samples: Record[]; datasetName?: string | null; hfToken?: string | null; + modelName?: string | null; + modelType?: "text" | "vision" | "audio" | "embeddings" | null; }; export type AiAssistMappingResponse = { @@ -88,6 +90,8 @@ export async function aiAssistMapping({ samples, datasetName, hfToken, + modelName, + modelType, }: AiAssistMappingArgs): Promise { const res = await authFetch("/api/datasets/ai-assist-mapping", { method: "POST", @@ -97,6 +101,8 @@ export async function aiAssistMapping({ samples: samples.slice(0, 5), dataset_name: datasetName || undefined, hf_token: hfToken || undefined, + model_name: modelName || undefined, + model_type: modelType || undefined, }), });