From 202780c32cd56799f9babc4800fb3a3b5dd07fc8 Mon Sep 17 00:00:00 2001 From: Roland Tannous Date: Tue, 10 Mar 2026 15:39:56 +0000 Subject: [PATCH] =?UTF-8?q?feat:=20Dataset=20Conversion=20Advisor=20?= =?UTF-8?q?=E2=80=94=20multi-pass=20LLM=20for=20non-conversational=20datas?= =?UTF-8?q?ets?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Non-conversational HF datasets (e.g. stanfordnlp/snli) were naively mapped column→role, producing poor training results. The AI Assist button now runs a 3-pass advisor using Qwen 7B that: 1. Fetches the HF dataset card/README to understand the dataset purpose 2. Classifies the dataset type and determines if conversion is needed 3. Generates a system prompt, user/assistant templates with {column} placeholders, and label mappings (e.g. 0→entailment) 4. Validates the conversion quality (score ≥7/10 required) Architecture: advisor metadata flows as __-prefixed keys in custom_format_mapping (e.g. __system_prompt, __user_template, __assistant_template, __label_mapping). The existing _apply_user_mapping() detects these keys and routes to template-based conversation construction. No __ keys = existing simple mode (backwards compatible). Backend: upgraded llm_assist.py (7B default, multi-pass advisor, HF card fetching), extended API models, added _apply_template_mapping() to dataset_utils.py. Frontend: extended store with advisor state fields, wired AI Assist to store templates/system prompt, inject __ metadata in training request, show advisor notification banner in mapping card. --- studio/backend/models/datasets.py | 11 +- studio/backend/models/training.py | 9 +- studio/backend/routes/datasets.py | 43 +- .../backend/utils/datasets/dataset_utils.py | 89 +++- studio/backend/utils/datasets/llm_assist.py | 441 +++++++++++++++++- .../dataset-preview-dialog-mapping.tsx | 10 +- .../sections/dataset-preview-dialog.tsx | 22 +- .../src/features/training/api/datasets-api.ts | 11 + .../src/features/training/api/mappers.ts | 20 +- .../training/stores/training-config-store.ts | 35 +- .../src/features/training/types/api.ts | 2 +- .../src/features/training/types/config.ts | 13 + 12 files changed, 672 insertions(+), 34 deletions(-) diff --git a/studio/backend/models/datasets.py b/studio/backend/models/datasets.py index d891d099b6..4aed5125dc 100644 --- a/studio/backend/models/datasets.py +++ b/studio/backend/models/datasets.py @@ -49,13 +49,22 @@ class AiAssistMappingRequest(BaseModel): columns: List[str] 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 class AiAssistMappingResponse(BaseModel): - """Response from LLM-assisted column classification.""" + """Response from LLM-assisted column classification and conversion advice.""" success: bool suggested_mapping: Optional[Dict[str, str]] = None warning: Optional[str] = None + # Conversion advisor fields + system_prompt: Optional[str] = None + user_template: Optional[str] = None + assistant_template: Optional[str] = None + label_mapping: Optional[Dict[str, Dict[str, str]]] = None + dataset_type: Optional[str] = None + is_conversational: Optional[bool] = None + user_notification: Optional[str] = None class UploadDatasetResponse(BaseModel): diff --git a/studio/backend/models/training.py b/studio/backend/models/training.py index c3b8d99bb4..73d42e464f 100644 --- a/studio/backend/models/training.py +++ b/studio/backend/models/training.py @@ -39,9 +39,14 @@ class TrainingStartRequest(BaseModel): if isinstance(values, dict) and "split" in values: values.setdefault("train_split", values.pop("split")) return values - custom_format_mapping: Optional[Dict[str, str]] = Field( + custom_format_mapping: Optional[Dict[str, Any]] = Field( None, - description="User-provided column-to-role mapping, e.g. {'image': 'image', 'caption': 'text'} for VLM or {'instruction': 'user', 'output': 'assistant'} for LLM" + description=( + "User-provided column-to-role mapping, e.g. {'image': 'image', 'caption': 'text'} " + "for VLM or {'instruction': 'user', 'output': 'assistant'} for LLM. " + "Enhanced format includes __system_prompt, __user_template, " + "__assistant_template, __label_mapping metadata keys." + ), ) # Training parameters num_epochs: int = Field(1, description="Number of training epochs") diff --git a/studio/backend/routes/datasets.py b/studio/backend/routes/datasets.py index 318eef561c..8deb1b1978 100644 --- a/studio/backend/routes/datasets.py +++ b/studio/backend/routes/datasets.py @@ -489,14 +489,17 @@ def ai_assist_mapping( current_subject: str = Depends(get_current_subject), ): """ - Run LLM-assisted column classification on demand (user-triggered). + Run LLM-assisted dataset conversion advisor (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. + Multi-pass analysis using a 7B helper model: + Pass 1: Classify dataset type from HF card + samples + Pass 2: Generate conversion strategy (system prompt, templates) + Pass 3: Validate conversion quality + + Falls back to simple column classification if the advisor fails. """ try: - from utils.datasets.llm_assist import llm_classify_columns, llm_generate_dataset_warning + from utils.datasets.llm_assist import llm_conversion_advisor # Truncate sample values for the LLM prompt truncated = [ @@ -504,31 +507,29 @@ def ai_assist_mapping( for s in request.samples[:5] ] - mapping = llm_classify_columns( + result = llm_conversion_advisor( column_names=request.columns, samples=truncated, + dataset_name=request.dataset_name, + hf_token=request.hf_token, ) - if mapping: - # Keep only conversation roles, not metadata - conversation_mapping = { - col: role for col, role in mapping.items() - if role in ("user", "assistant", "system") - } + + if result and result.get("success"): return AiAssistMappingResponse( success=True, - suggested_mapping=conversation_mapping, + suggested_mapping=result.get("suggested_mapping"), + system_prompt=result.get("system_prompt"), + user_template=result.get("user_template"), + assistant_template=result.get("assistant_template"), + label_mapping=result.get("label_mapping"), + dataset_type=result.get("dataset_type"), + is_conversational=result.get("is_conversational"), + user_notification=result.get("user_notification"), ) - # 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.", + warning="AI could not determine column roles. Please assign them manually.", ) except Exception as e: diff --git a/studio/backend/utils/datasets/dataset_utils.py b/studio/backend/utils/datasets/dataset_utils.py index 5de54a387d..1d1e2e78c0 100644 --- a/studio/backend/utils/datasets/dataset_utils.py +++ b/studio/backend/utils/datasets/dataset_utils.py @@ -196,12 +196,23 @@ def _apply_user_mapping(dataset, mapping: dict, batch_size: int = 1000): Accepts chatml (user/assistant/system), sharegpt (human/gpt/system), and alpaca (instruction/input/output) role names — all normalised to chatml output. + If the mapping contains ``__``-prefixed metadata keys (from the conversion + advisor), routes to template-based conversion instead of simple role mapping. + Returns: Dataset with single 'conversations' column """ + # Split metadata from column roles + meta = {k: v for k, v in mapping.items() if k.startswith("__")} + column_roles = {k: v for k, v in mapping.items() if not k.startswith("__")} + + if meta: + return _apply_template_mapping(dataset, column_roles, meta, batch_size) + + # ── Simple mode (original logic) ── # Pre-compute: group columns by canonical chatml role role_groups: dict[str, list[str]] = {r: [] for r in _CHATML_ROLE_ORDER} - for col_name, role in mapping.items(): + for col_name, role in column_roles.items(): canonical = _TO_CHATML.get(role) if canonical: role_groups[canonical].append(col_name) @@ -222,6 +233,82 @@ def _apply_user_mapping(dataset, mapping: dict, batch_size: int = 1000): return dataset.map(_convert, batched=True, batch_size=batch_size, remove_columns=dataset.column_names) +def _apply_template_mapping( + dataset, column_roles: dict, meta: dict, batch_size: int = 1000 +): + """ + Apply template-based mapping for non-conversational datasets. + + Uses ``__system_prompt``, ``__user_template``, ``__assistant_template``, + and ``__label_mapping`` metadata keys to construct conversations with + proper context and formatting. + + Returns: + Dataset with single 'conversations' column + """ + system_prompt = meta.get("__system_prompt", "") + user_template = meta.get("__user_template", "") + assistant_template = meta.get("__assistant_template", "") + label_mapping = meta.get("__label_mapping", {}) # {col: {int_str: label_str}} + + all_columns = list(dataset.column_names) + import logging as _log + _log.getLogger(__name__).info( + f"Applying template mapping: sys={bool(system_prompt)}, " + f"user_tpl={bool(user_template)}, asst_tpl={bool(assistant_template)}, " + f"label_map={list(label_mapping.keys())}" + ) + + def _convert(examples): + num = len(next(iter(examples.values()))) + conversations = [] + for i in range(num): + convo = [] + + # Build value dict for template interpolation + row_values = {} + for col in all_columns: + val = examples[col][i] + str_val = str(val) if val is not None else "" + + # Apply label mapping if this column has one + if col in label_mapping and isinstance(label_mapping[col], dict): + mapped = label_mapping[col].get(str_val, str_val) + row_values[col] = mapped + row_values[f"{col}_name"] = mapped + else: + row_values[col] = str_val + row_values[f"{col}_name"] = str_val + + # System prompt (static string, not from any column) + if system_prompt: + convo.append({"role": "system", "content": system_prompt}) + + # User message from template + if user_template: + try: + user_content = user_template.format(**row_values) + except (KeyError, IndexError): + user_content = user_template + convo.append({"role": "user", "content": user_content}) + + # Assistant message from template + if assistant_template: + try: + asst_content = assistant_template.format(**row_values) + except (KeyError, IndexError): + asst_content = assistant_template + convo.append({"role": "assistant", "content": asst_content}) + + conversations.append(convo) + return {"conversations": conversations} + + return dataset.map( + _convert, batched=True, batch_size=batch_size, + remove_columns=dataset.column_names, + ) + + def _apply_user_mapping_alpaca(dataset, mapping: dict, batch_size: int = 1000): """ Apply user-provided column mapping to convert dataset to Alpaca format. diff --git a/studio/backend/utils/datasets/llm_assist.py b/studio/backend/utils/datasets/llm_assist.py index d72d57232d..65df51ff15 100644 --- a/studio/backend/utils/datasets/llm_assist.py +++ b/studio/backend/utils/datasets/llm_assist.py @@ -16,14 +16,19 @@ Architecture: import json import logging import os +import re +import textwrap +import time from itertools import islice -from typing import Optional +from typing import Any, Optional logger = logging.getLogger(__name__) -DEFAULT_HELPER_MODEL_REPO = "Qwen/Qwen2.5-3B-Instruct-GGUF" +DEFAULT_HELPER_MODEL_REPO = "Qwen/Qwen2.5-7B-Instruct-GGUF" DEFAULT_HELPER_MODEL_VARIANT = "Q8_0" +README_MAX_CHARS = 1500 + def precache_helper_gguf(): """ @@ -325,3 +330,435 @@ def llm_generate_dataset_warning( print(f"🤖 LLM-generated warning: {warning}") return warning + + +# ─── Dataset Conversion Advisor ────────────────────────────────────── + + +def _parse_json_response(text: str) -> Optional[dict]: + """Parse JSON from LLM response, handling markdown fences and noise.""" + if not text: + return None + + cleaned = text.strip() + + # Strip markdown code fences + if cleaned.startswith("```"): + lines = cleaned.split("\n") + end = -1 if lines[-1].strip().startswith("```") else len(lines) + cleaned = "\n".join(lines[1:end]).strip() + + # Try direct parse + try: + obj = json.loads(cleaned) + if isinstance(obj, dict): + return obj + except json.JSONDecodeError: + pass + + # Greedy match for outermost {...} + match = re.search(r"\{.*\}", cleaned, re.DOTALL) + if match: + try: + obj = json.loads(match.group()) + if isinstance(obj, dict): + return obj + except json.JSONDecodeError: + pass + + return None + + +def _generate_with_backend( + backend, messages: list[dict], max_tokens: int = 512 +) -> str: + """Run one chat completion on an already-loaded backend. Returns raw text.""" + cumulative = "" + for text in backend.generate_chat_completion( + messages=messages, + temperature=0.1, + top_p=0.9, + top_k=20, + max_tokens=max_tokens, + repetition_penalty=1.0, + ): + cumulative = text + return cumulative.strip() + + +def fetch_hf_dataset_card( + dataset_name: str, hf_token: Optional[str] = None +) -> tuple[Optional[str], Optional[dict]]: + """ + Fetch HF dataset card (README) and metadata. + + Returns: + (readme_text, metadata_dict) or (None, None) on failure. + """ + try: + from huggingface_hub import DatasetCard + + card = DatasetCard.load(dataset_name, token=hf_token) + readme = card.text or "" + + # Truncate at sentence boundary + if len(readme) > README_MAX_CHARS: + cut = readme[:README_MAX_CHARS].rfind(".") + if cut > README_MAX_CHARS // 2: + readme = readme[: cut + 1] + "\n[...truncated]" + else: + readme = readme[:README_MAX_CHARS] + "\n[...truncated]" + + # Extract metadata from YAML frontmatter + metadata = {} + if card.data: + for key in ( + "task_categories", "task_ids", "language", + "size_categories", "tags", "license", "pretty_name", + ): + val = getattr(card.data, key, None) + if val is not None: + metadata[key] = val + + logger.info(f"Fetched dataset card: {len(readme)} chars, {len(metadata)} metadata fields") + return readme, metadata + + except Exception as e: + logger.warning(f"Could not fetch dataset card for {dataset_name}: {e}") + return None, None + + +def _run_multi_pass_advisor( + columns: list[str], + samples: list[dict], + dataset_name: Optional[str] = None, + dataset_card: Optional[str] = None, + dataset_metadata: Optional[dict] = None, +) -> Optional[dict[str, Any]]: + """ + Multi-pass LLM analysis: classify → convert → validate. + + Keeps model loaded across all passes. Returns combined result dict or None. + """ + if os.environ.get("UNSLOTH_HELPER_MODEL_DISABLE", "").strip() in ("1", "true"): + return None + + repo = os.environ.get("UNSLOTH_HELPER_MODEL_REPO", DEFAULT_HELPER_MODEL_REPO) + variant = os.environ.get("UNSLOTH_HELPER_MODEL_VARIANT", DEFAULT_HELPER_MODEL_VARIANT) + + backend = None + try: + from core.inference.llama_cpp import LlamaCppBackend + + backend = LlamaCppBackend() + print(f"🤖 Loading advisor model: {repo} ({variant})...") + t0 = time.monotonic() + + ok = backend.load_model( + hf_repo=repo, + hf_variant=variant, + model_identifier=f"advisor:{repo}:{variant}", + is_vision=False, + n_ctx=2048, + n_gpu_layers=-1, + ) + if not ok: + logger.warning("Advisor model failed to start") + return None + + print(f"🤖 Advisor model loaded in {time.monotonic() - t0:.1f}s") + + # ── Format samples ── + samples_text = "" + for i, row in enumerate(samples[:5], 1): + parts = [f" {col}: {str(row.get(col, ''))[:200]}" for col in columns] + samples_text += f"Row {i}:\n" + "\n".join(parts) + "\n" + + metadata_str = ( + json.dumps(dataset_metadata, indent=2, default=str)[:500] + if dataset_metadata else "N/A" + ) + card_excerpt = (dataset_card or "")[:1200] or "N/A" + + # ── Pass 1: Classify ── + print("🤖 Pass 1: Classifying dataset...", flush=True) + t1 = time.monotonic() + messages1 = [ + { + "role": "system", + "content": ( + "You are a dataset analyst specializing in HuggingFace datasets for LLM fine-tuning. " + "You classify datasets and determine if they can be used directly for conversational " + "fine-tuning or if they need conversion. Respond with ONLY valid JSON, no explanation." + ), + }, + { + "role": "user", + "content": textwrap.dedent(f"""\ + Analyze this HuggingFace dataset and classify it. + + DATASET CARD (excerpt): + {card_excerpt} + + METADATA: + {metadata_str} + + COLUMNS: {columns} + + SAMPLE DATA: + {samples_text} + + Respond with a JSON object: + {{ + "dataset_type": "", + "is_conversational": , + "needs_conversion": , + "description": "<1-2 sentence description of what this dataset is for>", + "task_description": "" + }}"""), + }, + ] + raw1 = _generate_with_backend(backend, messages1, max_tokens=256) + pass1 = _parse_json_response(raw1) + print(f"🤖 Pass 1 done ({time.monotonic() - t1:.1f}s): {pass1}", flush=True) + + if not pass1: + logger.warning(f"Advisor Pass 1 failed to produce JSON: {raw1[:200]}") + return None + + # If dataset is already conversational, skip passes 2-3 + if pass1.get("is_conversational") and not pass1.get("needs_conversion"): + return { + "success": True, + "dataset_type": pass1.get("dataset_type"), + "is_conversational": True, + "user_notification": ( + "This dataset is already in conversational format. " + "No conversion needed — columns can be mapped directly." + ), + } + + # ── Pass 2: Conversion strategy ── + print("🤖 Pass 2: Generating conversion strategy...", flush=True) + t2 = time.monotonic() + messages2 = [ + { + "role": "system", + "content": ( + "You are a dataset conversion specialist for LLM fine-tuning. " + "You design strategies to convert non-conversational datasets into " + "user/assistant conversation format. Respond with ONLY valid JSON." + ), + }, + { + "role": "user", + "content": textwrap.dedent(f"""\ + This dataset was classified as: + {json.dumps(pass1, indent=2)} + + COLUMNS: {columns} + + SAMPLE DATA: + {samples_text} + + Design a conversion strategy to turn this into conversation format for fine-tuning. + The strategy should create a system prompt, a user message template, and an assistant message template. + + For the user template, use {{column_name}} placeholders for column values. + For the assistant template, use {{column_name}} placeholders. + If a column has integer values that represent categories, provide a label mapping. + + Respond with a JSON object: + {{ + "system_prompt": "", + "user_template": "