diff --git a/studio/backend/core/training/trainer.py b/studio/backend/core/training/trainer.py
index 10f831c51f..b62c1bd35d 100644
--- a/studio/backend/core/training/trainer.py
+++ b/studio/backend/core/training/trainer.py
@@ -99,6 +99,71 @@ class UnslothTrainer:
'is_lora': True, # Default to LoRA
}
+ def pre_detect_and_load_tokenizer(
+ self,
+ model_name: str,
+ max_seq_length: int = 2048,
+ hf_token: Optional[str] = None,
+ is_dataset_image: bool = False,
+ is_dataset_audio: bool = False,
+ trust_remote_code: bool = False,
+ ) -> None:
+ """Lightweight detection and tokenizer load — no model weights, no VRAM.
+
+ Sets is_vlm, _audio_type, is_audio_vlm, model_name and loads a
+ lightweight tokenizer for dataset formatting. Call this before
+ load_and_format_dataset() when you want to process the dataset
+ BEFORE loading the training model (avoids VRAM contention with
+ the LLM-assisted detection helper).
+
+ load_model() may be called afterwards — it will re-detect and load
+ the full model + tokenizer, overwriting the lightweight one set here.
+ """
+ self.model_name = model_name
+ self.max_seq_length = max_seq_length
+ self.trust_remote_code = trust_remote_code
+
+ if hf_token:
+ os.environ["HF_TOKEN"] = hf_token
+
+ # --- Detect audio type (reads config.json only, no VRAM) ---
+ self._audio_type = detect_audio_type(model_name, hf_token)
+ if self._audio_type == 'audio_vlm':
+ self.is_audio = False
+ self.is_audio_vlm = is_dataset_audio
+ self._audio_type = None
+ else:
+ self.is_audio = self._audio_type is not None
+ self.is_audio_vlm = False
+
+ if not self.is_audio and not self.is_audio_vlm:
+ self._cuda_audio_used = False
+
+ # --- Detect VLM ---
+ vision = is_vision_model(model_name) if not self.is_audio else False
+ self.is_vlm = not self.is_audio_vlm and vision and is_dataset_image
+
+ logger.info(
+ "pre_detect: audio_type=%s, is_audio=%s, is_audio_vlm=%s, is_vlm=%s",
+ self._audio_type, self.is_audio, self.is_audio_vlm, self.is_vlm,
+ )
+
+ # --- Load lightweight tokenizer/processor (CPU only, no VRAM) ---
+ # Whisper needs AutoProcessor (has feature_extractor + tokenizer).
+ # All others work with AutoTokenizer (CSM loads its own processor inline).
+ if self._audio_type == 'whisper':
+ from transformers import AutoProcessor
+ self.tokenizer = AutoProcessor.from_pretrained(
+ model_name, trust_remote_code=trust_remote_code, token=hf_token,
+ )
+ else:
+ from transformers import AutoTokenizer
+ self.tokenizer = AutoTokenizer.from_pretrained(
+ model_name, trust_remote_code=trust_remote_code, token=hf_token,
+ )
+
+ logger.info("Pre-loaded tokenizer for %s", model_name)
+
def add_progress_callback(self, callback: Callable[[TrainingProgress], None]):
"""Add callback for training progress updates"""
self.progress_callbacks.append(callback)
@@ -2653,6 +2718,16 @@ class UnslothTrainer:
)
logger.info(f"Post-filter dataset size: {filtered_len} samples\n")
+ # [DEBUG] Decode first sample AFTER train_on_completions applied
+ try:
+ _row = self.trainer.train_dataset[0]
+ _space = self.tokenizer(" ", add_special_tokens=False).input_ids[0]
+ print("[DEBUG] === After train_on_completions ===", flush=True)
+ print(f"[DEBUG] input_ids decoded:\n{self.tokenizer.decode(_row['input_ids'])}\n", flush=True)
+ print(f"[DEBUG] labels decoded (-100 → space):\n{self.tokenizer.decode([_space if x == -100 else x for x in _row['labels']])}\n", flush=True)
+ except Exception as _dbg_e:
+ print(f"[DEBUG] Could not decode post-completions sample: {_dbg_e}", flush=True)
+
except Exception as e:
logger.warning(f"Failed to apply train on responses only: {e}")
train_on_responses_enabled = False
diff --git a/studio/backend/core/training/worker.py b/studio/backend/core/training/worker.py
index 8083bb9483..90505d7596 100644
--- a/studio/backend/core/training/worker.py
+++ b/studio/backend/core/training/worker.py
@@ -212,11 +212,87 @@ def run_training_process(
stop_thread.start()
# ── 4. Execute the training pipeline ──
+ # Order: detect → dataset → model → prepare → train
+ # Dataset processing (including LLM-assisted detection) runs BEFORE model
+ # loading so both never occupy VRAM at the same time.
try:
hf_token = config.get("hf_token", "")
hf_token = hf_token if hf_token and hf_token.strip() else None
- # Load model
+ # ── 4a. Lightweight detection + tokenizer (no VRAM) ──
+ _send_status(event_queue, "Detecting model type...")
+ trainer.pre_detect_and_load_tokenizer(
+ model_name=model_name,
+ max_seq_length=config["max_seq_length"],
+ hf_token=hf_token,
+ is_dataset_image=config.get("is_dataset_image", False),
+ is_dataset_audio=config.get("is_dataset_audio", False),
+ trust_remote_code=config.get("trust_remote_code", False),
+ )
+ if trainer.should_stop:
+ event_queue.put({"type": "complete", "output_dir": None, "ts": time.time()})
+ return
+
+ # ── 4b. Load and format dataset (LLM helper may use VRAM briefly) ──
+ _send_status(event_queue, "Loading and formatting dataset...")
+ hf_dataset = config.get("hf_dataset", "")
+ dataset_result = trainer.load_and_format_dataset(
+ dataset_source=hf_dataset if hf_dataset and hf_dataset.strip() else None,
+ format_type=config.get("format_type", ""),
+ local_datasets=config.get("local_datasets") or None,
+ custom_format_mapping=config.get("custom_format_mapping"),
+ subset=config.get("subset"),
+ train_split=config.get("train_split", "train"),
+ eval_split=config.get("eval_split"),
+ eval_steps=config.get("eval_steps", 0.00),
+ dataset_slice_start=config.get("dataset_slice_start"),
+ dataset_slice_end=config.get("dataset_slice_end"),
+ )
+
+ if isinstance(dataset_result, tuple):
+ dataset, eval_dataset = dataset_result
+ else:
+ dataset = dataset_result
+ eval_dataset = None
+
+ # [DEBUG] Print first sample before model is loaded
+ # dataset is a dict {"dataset": , "detected_format": ..., ...}
+ # or a raw Dataset for audio paths
+ try:
+ ds = dataset["dataset"] if isinstance(dataset, dict) else dataset
+ print(f"\n[DEBUG] Dataset loaded BEFORE model. type={type(ds).__name__}, len={len(ds)}", flush=True)
+ print(f"[DEBUG] Columns: {ds.column_names}", flush=True)
+ sample = ds[0]
+ preview = {k: str(v)[:300] for k, v in sample.items()}
+ print(f"[DEBUG] First sample: {preview}\n", flush=True)
+ except Exception as e:
+ print(f"[DEBUG] Could not preview first sample: {type(e).__name__}: {e}", flush=True)
+
+ # Disable eval if eval_steps <= 0
+ eval_steps = config.get("eval_steps", 0.00)
+ if eval_steps is not None and float(eval_steps) <= 0:
+ eval_dataset = None
+
+ # Tell the parent process that eval is configured so the frontend
+ # shows "Waiting for first evaluation step..." instead of "not configured"
+ if eval_dataset is not None:
+ event_queue.put({
+ "type": "eval_configured",
+ "ts": time.time(),
+ })
+
+ if dataset is None or trainer.should_stop:
+ if trainer.should_stop:
+ event_queue.put({"type": "complete", "output_dir": None, "ts": time.time()})
+ else:
+ event_queue.put({
+ "type": "error",
+ "error": trainer.training_progress.error or "Failed to load dataset",
+ "stack": "", "ts": time.time(),
+ })
+ return
+
+ # ── 4c. Load training model (uses VRAM — dataset already formatted) ──
_send_status(event_queue, "Loading model...")
success = trainer.load_model(
model_name=model_name,
@@ -239,7 +315,7 @@ def run_training_process(
})
return
- # Prepare model (LoRA or full finetuning)
+ # ── 4d. Prepare model (LoRA or full finetuning) ──
training_type = config.get("training_type", "LoRA/QLoRA")
use_lora = (training_type == "LoRA/QLoRA")
if use_lora:
@@ -273,52 +349,6 @@ def run_training_process(
})
return
- # Load dataset
- _send_status(event_queue, "Loading and formatting dataset...")
- hf_dataset = config.get("hf_dataset", "")
- dataset_result = trainer.load_and_format_dataset(
- dataset_source=hf_dataset if hf_dataset and hf_dataset.strip() else None,
- format_type=config.get("format_type", ""),
- local_datasets=config.get("local_datasets") or None,
- custom_format_mapping=config.get("custom_format_mapping"),
- subset=config.get("subset"),
- train_split=config.get("train_split", "train"),
- eval_split=config.get("eval_split"),
- eval_steps=config.get("eval_steps", 0.00),
- dataset_slice_start=config.get("dataset_slice_start"),
- dataset_slice_end=config.get("dataset_slice_end"),
- )
-
- if isinstance(dataset_result, tuple):
- dataset, eval_dataset = dataset_result
- else:
- dataset = dataset_result
- eval_dataset = None
-
- # Disable eval if eval_steps <= 0
- eval_steps = config.get("eval_steps", 0.00)
- if eval_steps is not None and float(eval_steps) <= 0:
- eval_dataset = None
-
- # Tell the parent process that eval is configured so the frontend
- # shows "Waiting for first evaluation step..." instead of "not configured"
- if eval_dataset is not None:
- event_queue.put({
- "type": "eval_configured",
- "ts": time.time(),
- })
-
- if dataset is None or trainer.should_stop:
- if trainer.should_stop:
- event_queue.put({"type": "complete", "output_dir": None, "ts": time.time()})
- else:
- event_queue.put({
- "type": "error",
- "error": trainer.training_progress.error or "Failed to load dataset",
- "stack": "", "ts": time.time(),
- })
- return
-
# Convert learning rate
try:
lr_value = float(config.get("learning_rate", "2e-4"))
diff --git a/studio/backend/main.py b/studio/backend/main.py
index a5cb337d69..d75c76ac48 100644
--- a/studio/backend/main.py
+++ b/studio/backend/main.py
@@ -72,6 +72,17 @@ async def lifespan(app: FastAPI):
f"GPU sm_{sm_version} detected — setting UNSLOTH_FLEX_ATTENTION=0"
)
+ # Pre-cache the helper GGUF model for LLM-assisted dataset detection.
+ # Runs in a background thread so it doesn't block server startup.
+ import threading
+ def _precache():
+ try:
+ from utils.datasets.llm_assist import precache_helper_gguf
+ precache_helper_gguf()
+ except Exception:
+ pass # non-critical
+ threading.Thread(target=_precache, daemon=True).start()
+
if not storage.is_initialized():
setup_token = secrets.token_urlsafe(32)
storage.save_setup_token(setup_token)
diff --git a/studio/backend/models/datasets.py b/studio/backend/models/datasets.py
index 73c75dc650..f84c98fc0a 100644
--- a/studio/backend/models/datasets.py
+++ b/studio/backend/models/datasets.py
@@ -44,6 +44,29 @@ 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
+ hf_token: Optional[str] = None # For fetching dataset card
+ model_name: Optional[str] = None
+ model_type: Optional[str] = None
+
+
+class AiAssistMappingResponse(BaseModel):
+ """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
+ 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):
"""Response with stored dataset path for training."""
filename: str = Field(..., description="Original filename")
diff --git a/studio/backend/models/training.py b/studio/backend/models/training.py
index c2724a73b1..8773c75941 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 e0029d8671..29f7cba005 100644
--- a/studio/backend/routes/datasets.py
+++ b/studio/backend/routes/datasets.py
@@ -30,6 +30,8 @@ logger = get_logger(__name__)
from models.datasets import (
+ AiAssistMappingRequest,
+ AiAssistMappingResponse,
CheckFormatRequest,
CheckFormatResponse,
LocalDatasetItem,
@@ -432,18 +434,19 @@ def check_format(
else:
preview_samples = _serialize_preview_rows(preview_slice)
- # Lightweight URL-based image detection for VLM datasets
- warning = None
+ # Collect warnings: from check_dataset_format + URL-based image detection
+ warning = result.get("warning")
image_col = result.get("detected_image_column")
if image_col and image_col in (result.get("columns") or []):
try:
sample_val = preview_slice[0][image_col]
if isinstance(sample_val, str) and sample_val.startswith(("http://", "https://")):
- warning = (
+ url_warning = (
"This dataset contains image URLs instead of embedded images. "
"Images will be downloaded during training, which may be slow for large datasets."
)
logger.info(f"URL-based image column detected: {image_col}")
+ warning = f"{warning} {url_warning}" if warning else url_warning
except Exception:
pass
@@ -472,3 +475,62 @@ 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 dataset conversion advisor (user-triggered).
+
+ 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_conversion_advisor
+
+ # 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]
+ ]
+
+ result = llm_conversion_advisor(
+ column_names=request.columns,
+ 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"):
+ return AiAssistMappingResponse(
+ success=True,
+ 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"),
+ )
+
+ return AiAssistMappingResponse(
+ success=False,
+ warning="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 e1f45610ab..40b0b09353 100644
--- a/studio/backend/utils/datasets/dataset_utils.py
+++ b/studio/backend/utils/datasets/dataset_utils.py
@@ -18,6 +18,8 @@ All internal utilities have been moved to separate modules:
- model_mappings: TEMPLATE_TO_MODEL_MAPPER
"""
+import json
+
# Import from modular files
from .format_detection import (
detect_dataset_format,
@@ -89,6 +91,21 @@ def check_dataset_format(dataset, is_vlm: bool = False) -> dict:
vlm_structure = detect_vlm_dataset_structure(dataset)
requires_mapping = vlm_structure["format"] == "unknown"
+ warning = None
+ if requires_mapping:
+ img_col = vlm_structure.get("image_column")
+ txt_col = vlm_structure.get("text_column")
+ missing = []
+ if not img_col:
+ missing.append("image")
+ if not txt_col:
+ missing.append("text")
+ if missing:
+ warning = (
+ f"Could not auto-detect {' or '.join(missing)} column. "
+ "Please assign image and text columns manually."
+ )
+
return {
"requires_manual_mapping": requires_mapping,
"detected_format": vlm_structure["format"],
@@ -98,6 +115,7 @@ def check_dataset_format(dataset, is_vlm: bool = False) -> dict:
"detected_text_column": vlm_structure.get("text_column"),
"is_image": multimodal_info["is_image"],
"multimodal_columns": multimodal_info.get("multimodal_columns"),
+ "warning": warning,
**audio_fields,
}
@@ -118,7 +136,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
@@ -137,6 +155,7 @@ def check_dataset_format(dataset, is_vlm: bool = False) -> dict:
**audio_fields,
}
else:
+ # Heuristic failed — user must map manually (or use AI Assist)
return {
"requires_manual_mapping": True,
"detected_format": "unknown",
@@ -146,6 +165,10 @@ def check_dataset_format(dataset, is_vlm: bool = False) -> dict:
"detected_text_column": None,
"is_image": False,
"multimodal_columns": None,
+ "warning": (
+ f"Could not auto-detect column roles for columns: {columns}. "
+ "Please assign roles manually, or use AI Assist."
+ ),
**audio_fields,
}
@@ -179,12 +202,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)
@@ -205,6 +239,98 @@ 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 _extract_column_value(val, col: str, label_mapping: dict) -> str:
+ """Extract a string value from a column, handling complex types and label mapping."""
+ # Handle complex types (dicts, lists) — extract useful text instead of raw repr
+ if isinstance(val, dict):
+ # Common pattern: {"text": [...]} in QA datasets
+ if "text" in val:
+ inner = val["text"]
+ str_val = inner[0] if isinstance(inner, list) and inner else str(inner)
+ else:
+ str_val = json.dumps(val, ensure_ascii=False)
+ elif isinstance(val, list):
+ str_val = val[0] if len(val) == 1 else ", ".join(str(v) for v in val)
+ else:
+ 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):
+ str_val = label_mapping[col].get(str_val, str_val)
+
+ return str_val
+
+
+def _apply_template_mapping(
+ dataset, column_roles: dict, meta: dict, batch_size: int = 1000
+):
+ """
+ Apply advisor-driven mapping for non-conversational datasets.
+
+ Groups columns by their assigned role (user/assistant), concatenates
+ values within each role into a single message, and injects an optional
+ system prompt. Label mapping is applied to convert integer labels
+ to human-readable strings.
+
+ Returns:
+ Dataset with single 'conversations' column
+ """
+ system_prompt = meta.get("__system_prompt", "")
+ label_mapping = meta.get("__label_mapping", {}) # {col: {int_str: label_str}}
+
+ # Group columns by canonical chatml role
+ role_groups: dict[str, list[str]] = {"user": [], "assistant": []}
+ for col, role in column_roles.items():
+ canonical = _TO_CHATML.get(role, role)
+ if canonical in role_groups:
+ role_groups[canonical].append(col)
+
+ import logging as _log
+ _log.getLogger(__name__).info(
+ f"Applying role mapping: sys={bool(system_prompt)}, "
+ f"user_cols={role_groups['user']}, asst_cols={role_groups['assistant']}, "
+ f"label_map={list(label_mapping.keys())}"
+ )
+
+ def _convert(examples):
+ num = len(next(iter(examples.values())))
+ conversations = []
+ for i in range(num):
+ convo = []
+
+ # System prompt (generated, static across all rows)
+ if system_prompt:
+ convo.append({"role": "system", "content": system_prompt})
+
+ # User message: concatenate all user-role column values
+ user_parts = []
+ for col in role_groups["user"]:
+ if col in examples:
+ user_parts.append(
+ _extract_column_value(examples[col][i], col, label_mapping)
+ )
+ if user_parts:
+ convo.append({"role": "user", "content": "\n".join(user_parts)})
+
+ # Assistant message: concatenate all assistant-role column values
+ asst_parts = []
+ for col in role_groups["assistant"]:
+ if col in examples:
+ asst_parts.append(
+ _extract_column_value(examples[col][i], col, label_mapping)
+ )
+ if asst_parts:
+ convo.append({"role": "assistant", "content": "\n".join(asst_parts)})
+
+ 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.
@@ -776,8 +902,22 @@ def format_and_template_dataset(
vlm_image_column = vlm_structure["image_column"]
if vlm_text_column is None or vlm_image_column is None:
+ columns = list(next(iter(dataset)).keys()) if dataset else []
+ issues = [
+ f"Could not auto-detect image and text columns from: {columns}",
+ f"VLM structure detected: {vlm_structure.get('format', 'unknown')}",
+ ]
+ friendly = None
+ try:
+ from .llm_assist import llm_generate_dataset_warning
+ friendly = llm_generate_dataset_warning(
+ issues, dataset_name=dataset_name, modality="vision",
+ column_names=columns,
+ )
+ except Exception:
+ pass
errors.append(
- f"Could not auto-detect image/text columns. Found: {vlm_structure}. "
+ friendly or f"Could not auto-detect image/text columns. Found: {vlm_structure}. "
)
return {
"dataset": dataset,
diff --git a/studio/backend/utils/datasets/format_conversion.py b/studio/backend/utils/datasets/format_conversion.py
index 624a10ba9b..fea2589489 100644
--- a/studio/backend/utils/datasets/format_conversion.py
+++ b/studio/backend/utils/datasets/format_conversion.py
@@ -435,7 +435,21 @@ def convert_to_vlm_format(
throughput = probe_total / probe_elapsed if probe_elapsed > 0 else 0
if fail_rate >= MAX_FAIL_RATE:
- msg = (
+ issues = [
+ f"{fail_rate:.0%} of the first {PROBE_SIZE} image URLs failed to download ({probe_fail}/{probe_total})",
+ "Images are external URLs, not embedded in the dataset",
+ ]
+ # Try LLM-friendly warning
+ friendly = None
+ try:
+ from .llm_assist import llm_generate_dataset_warning
+ friendly = llm_generate_dataset_warning(
+ issues, dataset_name=dataset_name, modality="vision",
+ column_names=[image_column, text_column],
+ )
+ except Exception:
+ pass
+ msg = friendly or (
f"⚠️ {fail_rate:.0%} of the first {PROBE_SIZE} images failed to download "
f"({probe_fail}/{probe_total}). "
"This dataset has too many broken or unreachable image URLs. "
@@ -524,7 +538,20 @@ def convert_to_vlm_format(
logger.info(f"⚠️ Skipped {failed_count}/{total} ({fail_rate:.0%}) samples with broken/unreachable images")
# For datasets that skipped the probe (small URL datasets), check fail rate now
if has_urls and fail_rate >= MAX_FAIL_RATE:
- msg = (
+ issues = [
+ f"{fail_rate:.0%} of images failed to download ({failed_count}/{total})",
+ "Images are external URLs, not embedded in the dataset",
+ ]
+ friendly = None
+ try:
+ from .llm_assist import llm_generate_dataset_warning
+ friendly = llm_generate_dataset_warning(
+ issues, dataset_name=dataset_name, modality="vision",
+ column_names=[image_column, text_column],
+ )
+ except Exception:
+ pass
+ msg = friendly or (
f"⚠️ {fail_rate:.0%} of images failed to download ({failed_count}/{total}). "
"This dataset has too many broken or unreachable image URLs. "
"Consider using a dataset with embedded images instead."
@@ -533,9 +560,25 @@ def convert_to_vlm_format(
raise ValueError(msg)
if len(converted_list) == 0:
+ issues = [
+ f"All {total} samples failed during VLM conversion — no usable images found",
+ f"Image column '{image_column}' may contain URLs that are no longer accessible, "
+ "or local file paths that don't exist",
+ ]
+ friendly = None
+ try:
+ from .llm_assist import llm_generate_dataset_warning
+ friendly = llm_generate_dataset_warning(
+ issues, dataset_name=dataset_name, modality="vision",
+ column_names=[image_column, text_column],
+ )
+ except Exception:
+ pass
raise ValueError(
- f"All {total} samples failed during VLM conversion — no usable images found. "
- "This dataset may contain only image URLs that are no longer accessible."
+ friendly or (
+ f"All {total} samples failed during VLM conversion — no usable images found. "
+ "This dataset may contain only image URLs that are no longer accessible."
+ )
)
logger.info(f"✅ Converted {len(converted_list)}/{total} samples")
diff --git a/studio/backend/utils/datasets/llm_assist.py b/studio/backend/utils/datasets/llm_assist.py
new file mode 100644
index 0000000000..f7016f5ff1
--- /dev/null
+++ b/studio/backend/utils/datasets/llm_assist.py
@@ -0,0 +1,837 @@
+# SPDX-License-Identifier: AGPL-3.0-only - See /studio/LICENSE.AGPL-3.0
+# Copyright © 2025 Unsloth AI
+
+"""
+LLM-assisted dataset analysis using an ephemeral GGUF helper model.
+
+Complements heuristic-based detection in format_detection.py and
+vlm_processing.py. Only invoked when heuristics are uncertain.
+
+Architecture:
+ - Instantiates LlamaCppBackend, loads model, runs completion(s), unloads.
+ - Not kept warm — VRAM is freed immediately after use.
+ - Gracefully degrades: returns None when unavailable (no binary, OOM, disabled).
+"""
+
+import json
+import logging
+import os
+import re
+import textwrap
+import time
+from itertools import islice
+from typing import Any, Optional
+
+logger = logging.getLogger(__name__)
+
+DEFAULT_HELPER_MODEL_REPO = "Qwen/Qwen2.5-7B-Instruct-GGUF"
+DEFAULT_HELPER_MODEL_VARIANT = "Q8_0"
+
+README_MAX_CHARS = 1500
+
+
+def precache_helper_gguf():
+ """
+ Pre-download the helper GGUF to HF cache.
+
+ Called on FastAPI startup in a background thread so subsequent
+ ``_run_with_helper()`` calls skip the download and only pay for
+ llama-server startup. No-op if already cached or disabled.
+ """
+ if os.environ.get("UNSLOTH_HELPER_MODEL_DISABLE", "").strip() in ("1", "true"):
+ return
+
+ repo = os.environ.get("UNSLOTH_HELPER_MODEL_REPO", DEFAULT_HELPER_MODEL_REPO)
+ variant = os.environ.get("UNSLOTH_HELPER_MODEL_VARIANT", DEFAULT_HELPER_MODEL_VARIANT)
+
+ try:
+ from huggingface_hub import HfApi, hf_hub_download
+
+ # Find the GGUF file matching the variant
+ api = HfApi()
+ files = api.list_repo_files(repo, repo_type="model")
+ gguf_files = [f for f in files if f.endswith(".gguf")]
+
+ # Find all GGUF files matching the variant (may be split into shards)
+ variant_lower = variant.lower().replace("-", "_")
+ matching = sorted(
+ f for f in gguf_files
+ if variant_lower in f.lower().replace("-", "_")
+ )
+
+ if matching:
+ logger.info(f"Pre-caching helper GGUF: {repo}/{matching[0]}"
+ + (f" (+{len(matching) - 1} shards)" if len(matching) > 1 else ""))
+ for target in matching:
+ hf_hub_download(repo_id=repo, filename=target)
+ logger.info(f"Helper GGUF cached: {len(matching)} file(s)")
+ else:
+ logger.warning(f"No GGUF matching variant '{variant}' in {repo}")
+ except Exception as e:
+ logger.warning(f"Failed to pre-cache helper GGUF: {e}")
+
+
+def _run_with_helper(prompt: str, max_tokens: int = 256) -> Optional[str]:
+ """
+ Load helper model, run one chat completion, unload.
+
+ Returns the completion text, or None on any failure.
+ """
+ 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()
+ logger.info(f"Loading helper model: {repo} ({variant})")
+ print(f"🤖 Loading helper model: {repo} ({variant})...")
+
+ ok = backend.load_model(
+ hf_repo=repo,
+ hf_variant=variant,
+ model_identifier=f"helper:{repo}:{variant}",
+ is_vision=False,
+ n_ctx=2048,
+ n_gpu_layers=-1,
+ )
+ if not ok:
+ logger.warning("Helper model failed to start")
+ return None
+
+ messages = [{"role": "user", "content": prompt}]
+ 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 # cumulative — last value is full text
+
+ result = cumulative.strip()
+ logger.info(f"Helper model response ({len(result)} chars)")
+ return result if result else None
+
+ except Exception as e:
+ logger.warning(f"Helper model failed: {e}")
+ return None
+
+ finally:
+ if backend is not None:
+ try:
+ backend.unload_model()
+ print("🤖 Helper model unloaded")
+ except Exception:
+ pass
+
+
+# ─── Public API ───────────────────────────────────────────────────────
+
+
+def llm_generate_vlm_instruction(
+ column_names: list[str],
+ samples: list[dict],
+ dataset_name: Optional[str] = None,
+) -> Optional[dict]:
+ """
+ Ask a helper LLM to generate a task-specific VLM instruction.
+
+ Called when heuristic instruction generation returns low confidence
+ or falls back to generic.
+
+ Args:
+ column_names: Column names in the dataset.
+ samples: 3-5 sample rows with text values (images replaced by "").
+ dataset_name: Optional HF dataset identifier for context.
+
+ Returns:
+ {"instruction": str, "confidence": 0.85} or None.
+ """
+ # Format samples for the prompt
+ formatted = ""
+ for i, row in enumerate(samples[:5], 1):
+ parts = []
+ for col in column_names:
+ val = str(row.get(col, ""))[:300]
+ parts.append(f" {col}: {val}")
+ formatted += f"Sample {i}:\n" + "\n".join(parts) + "\n\n"
+
+ prompt = (
+ "You are a dataset analyst. Given a vision-language dataset, generate ONE "
+ "instruction sentence that describes what the model should do with each image.\n\n"
+ f"Dataset: {dataset_name or 'unknown'}\n"
+ f"Columns: {column_names}\n\n"
+ f"{formatted}"
+ "Write ONE instruction sentence. Examples:\n"
+ '- "Solve the math problem shown in the image and explain your reasoning."\n'
+ '- "Transcribe all text visible in this image."\n'
+ '- "Answer the question about this image."\n\n'
+ "Respond with ONLY the instruction sentence, nothing else."
+ )
+
+ result = _run_with_helper(prompt, max_tokens=100)
+ if not result:
+ return None
+
+ # Clean up: strip quotes, ensure it's a single sentence
+ instruction = result.strip().strip('"').strip("'").strip()
+ # Reject obviously bad outputs (too short, too long, or multi-line)
+ if len(instruction) < 10 or len(instruction) > 200 or "\n" in instruction:
+ logger.warning(f"Helper model returned unusable instruction: {instruction!r}")
+ return None
+
+ print(f"🤖 LLM-generated instruction: {instruction}")
+ return {
+ "instruction": instruction,
+ "confidence": 0.85,
+ }
+
+
+def llm_classify_columns(
+ column_names: list[str],
+ samples: list[dict],
+) -> Optional[dict[str, str]]:
+ """
+ Ask a helper LLM to classify dataset columns into roles.
+
+ Called when heuristic column detection fails (returns None).
+
+ Args:
+ column_names: Column names in the dataset.
+ samples: 3-5 sample rows with values truncated to 200 chars.
+
+ Returns:
+ Dict mapping column_name → role ("user"|"assistant"|"system"|"metadata"),
+ or None on failure.
+ """
+ formatted = ""
+ for i, row in enumerate(samples[:5], 1):
+ parts = []
+ for col in column_names:
+ val = str(row.get(col, ""))[:200]
+ parts.append(f" {col}: {val}")
+ formatted += f"Sample {i}:\n" + "\n".join(parts) + "\n\n"
+
+ prompt = (
+ "Classify each column in this dataset into one of these roles:\n"
+ "- user: The input/question/prompt from the human\n"
+ "- assistant: The expected output/answer/response from the AI\n"
+ "- system: Context, persona, or task description\n"
+ "- metadata: IDs, scores, labels, timestamps — not part of conversation\n\n"
+ f"Columns: {column_names}\n\n"
+ f"{formatted}"
+ "Respond with ONLY a JSON object mapping column names to roles.\n"
+ 'Example: {"question": "user", "answer": "assistant", "id": "metadata"}'
+ )
+
+ result = _run_with_helper(prompt, max_tokens=200)
+ if not result:
+ return None
+
+ # Parse JSON from response (may have markdown fences)
+ text = result.strip()
+ if text.startswith("```"):
+ # Strip markdown code fence
+ lines = text.split("\n")
+ text = "\n".join(lines[1:-1] if lines[-1].strip() == "```" else lines[1:])
+ text = text.strip()
+
+ try:
+ mapping = json.loads(text)
+ except json.JSONDecodeError:
+ # Try to find JSON object in the response
+ import re
+ match = re.search(r"\{[^}]+\}", text)
+ if match:
+ try:
+ mapping = json.loads(match.group())
+ except json.JSONDecodeError:
+ logger.warning(f"Could not parse helper model JSON: {text!r}")
+ return None
+ else:
+ logger.warning(f"No JSON found in helper model response: {text!r}")
+ return None
+
+ if not isinstance(mapping, dict):
+ return None
+
+ # Validate: all values must be valid roles
+ valid_roles = {"user", "assistant", "system", "metadata"}
+ cleaned = {}
+ for col, role in mapping.items():
+ if col in column_names and isinstance(role, str) and role.lower() in valid_roles:
+ cleaned[col] = role.lower()
+
+ if not cleaned:
+ return None
+
+ # Must have at least user + assistant
+ roles_present = set(cleaned.values())
+ if "user" not in roles_present or "assistant" not in roles_present:
+ logger.warning(f"Helper model mapping missing user/assistant: {cleaned}")
+ return None
+
+ print(f"🤖 LLM-classified columns: {cleaned}")
+ return cleaned
+
+
+def llm_generate_dataset_warning(
+ issues: list[str],
+ dataset_name: Optional[str] = None,
+ modality: str = "text",
+ column_names: Optional[list[str]] = None,
+) -> Optional[str]:
+ """
+ Ask the helper LLM to turn technical dataset issues into a user-friendly warning.
+
+ Works for all modalities (text, vision, audio).
+
+ Args:
+ issues: List of technical issue descriptions found during analysis.
+ dataset_name: Optional HF dataset name.
+ modality: "text", "vision", or "audio".
+ column_names: Optional list of column names for context.
+
+ Returns:
+ A human-friendly warning string, or None on failure.
+ """
+ if not issues:
+ return None
+
+ issues_text = "\n".join(f"- {issue}" for issue in issues)
+ cols_text = f"\nColumns: {column_names}" if column_names else ""
+
+ prompt = (
+ "You are a helpful assistant. A user is trying to fine-tune a model on a dataset.\n"
+ "The following issues were found during dataset analysis:\n\n"
+ f"{issues_text}\n\n"
+ f"Dataset: {dataset_name or 'unknown'}\n"
+ f"Modality: {modality}"
+ f"{cols_text}\n\n"
+ "Write a brief, friendly explanation of what's wrong and what the user can do about it.\n"
+ "Keep it under 3 sentences. Be specific about the dataset."
+ )
+
+ result = _run_with_helper(prompt, max_tokens=200)
+ if not result:
+ return None
+
+ warning = result.strip()
+ # Reject obviously bad outputs
+ if len(warning) < 10 or len(warning) > 500:
+ return None
+
+ 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,
+ 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.
+
+ 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"
+
+ # ── 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()
+ messages1 = [
+ {
+ "role": "system",
+ "content": (
+ "You are a dataset analyst. Your job is to look at a HuggingFace dataset "
+ "and figure out what kind of data it contains and whether it is already in "
+ "a conversational format suitable for LLM fine-tuning. A dataset is "
+ '"conversational" if it already has columns like "messages", "conversations", '
+ 'or multiturn "user"/"assistant" pairs. Some datasets are NOT conversational '
+ "— 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}"
+ ),
+ },
+ {
+ "role": "user",
+ "content": textwrap.dedent(f"""\
+ Look at this HuggingFace dataset and classify it.
+
+ DATASET CARD (excerpt):
+ {card_excerpt}
+
+ METADATA:
+ {metadata_str}
+
+ COLUMNS: {columns}
+
+ SAMPLE DATA (first 3 rows):
+ {samples_text}
+
+ Based on the above, respond with this exact JSON structure:
+ {{
+ "dataset_type": "",
+ "is_conversational": ,
+ "needs_conversion": ,
+ "description": "",
+ "task_description": ""
+ }}
+
+ Respond with ONLY the JSON object. No markdown, no explanation."""),
+ },
+ ]
+ 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: Map columns to roles ──
+ print("🤖 Pass 2: Mapping columns to roles...", flush=True)
+ t2 = time.monotonic()
+ messages2 = [
+ {
+ "role": "system",
+ "content": (
+ "You are a data preparation assistant. Your job is to assign each column "
+ "in a dataset to a conversation role for LLM fine-tuning. There are exactly "
+ "two roles:\n"
+ '- "user" = This column contains INPUT that the model will receive as a prompt.\n'
+ '- "assistant" = This column contains OUTPUT that the model should learn to generate.\n\n'
+ "CRITICAL RULES:\n"
+ "1. There MUST be at least one column assigned to \"user\" AND at least one "
+ "column assigned to \"assistant\". Never assign all columns to the same role.\n"
+ "2. The column that contains the TARGET or OUTPUT or ANSWER or LABEL must "
+ "ALWAYS be assigned to \"assistant\". This is the thing the model should learn "
+ "to produce.\n"
+ "3. The columns that contain the SOURCE or INPUT or CONTEXT or QUESTION must "
+ "be assigned to \"user\". This is what the model receives.\n"
+ '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}"
+ ),
+ },
+ {
+ "role": "user",
+ "content": textwrap.dedent(f"""\
+ Here is a dataset that has been classified:
+
+ CLASSIFICATION:
+ {json.dumps(pass1, indent=2)}
+
+ COLUMNS AVAILABLE: {columns}
+
+ SAMPLE DATA (first 3 rows):
+ {samples_text}
+
+ Your task: assign each column to either "user", "assistant", or "skip".
+
+ Here are worked examples to guide you:
+
+ Example 1 — Summarization dataset with columns ["document", "summary"]:
+ "document" is the input text → "user"
+ "summary" is the output the model should generate → "assistant"
+ Result: {{"document": "user", "summary": "assistant"}}
+
+ Example 2 — Question answering dataset with columns ["context", "question", "answer"]:
+ "context" is input → "user"
+ "question" is input → "user"
+ "answer" is what the model should generate → "assistant"
+ Result: {{"context": "user", "question": "user", "answer": "assistant"}}
+
+ Example 3 — Classification dataset with columns ["text", "label"]:
+ "text" is input → "user"
+ "label" is the output the model should predict → "assistant"
+ Result: {{"text": "user", "label": "assistant"}}
+
+ Example 4 — Translation dataset with columns ["en", "fr"]:
+ "en" is the source language (input) → "user"
+ "fr" is the target language (output) → "assistant"
+ Result: {{"en": "user", "fr": "assistant"}}
+
+ Now apply this logic to the actual dataset columns listed above.
+
+ Respond with this exact JSON structure:
+ {{
+ "column_roles": {{
+ "": ""
+ }},
+ "label_mapping": ,
+ "notes": ""
+ }}
+
+ REMEMBER: There must be at least one "user" column AND at least one "assistant" column. If all columns are "user", you made a mistake — the output/target column should be "assistant".
+
+ Respond with ONLY the JSON object."""),
+ },
+ ]
+ raw2 = _generate_with_backend(backend, messages2, max_tokens=512)
+ pass2 = _parse_json_response(raw2)
+ print(f"🤖 Pass 2 done ({time.monotonic() - t2:.1f}s): {pass2}", flush=True)
+
+ if not pass2:
+ logger.warning(f"Advisor Pass 2 failed to produce JSON: {raw2[:200]}")
+ return None
+
+ # ── Extract and validate column roles from Pass 2 ──
+ column_roles = pass2.get("column_roles", {})
+ label_map = pass2.get("label_mapping") or {} # may be null
+
+ # Validate: must have at least one user AND one assistant
+ roles_present = set(column_roles.values())
+ if "user" not in roles_present or "assistant" not in roles_present:
+ print(
+ f"🤖 Pass 2 sanity fail: missing user or assistant role: {column_roles}",
+ flush=True,
+ )
+ return None # triggers fallback to simple classification
+
+ # ── Pass 3: System prompt (non-conversational datasets only) ──
+ sys_prompt = ""
+ dtype = pass1.get("dataset_type", "unknown")
+ is_conv = pass1.get("is_conversational", False)
+
+ if not is_conv:
+ print("🤖 Pass 3: Generating system prompt...", flush=True)
+ t3 = time.monotonic()
+
+ # Format label mapping info for the prompt
+ label_info = ""
+ if label_map:
+ for col, mapping in label_map.items():
+ if isinstance(mapping, dict) and mapping:
+ pairs = ", ".join(f"{k} = {v}" for k, v in mapping.items())
+ label_info += f"\nLabel mapping for '{col}': {pairs}"
+
+ # Describe the role assignments for context
+ user_cols = [c for c, r in column_roles.items() if r == "user"]
+ asst_cols = [c for c, r in column_roles.items() if r == "assistant"]
+ task_desc = pass1.get("task_description") or pass1.get("description", "")
+
+ messages3 = [
+ {
+ "role": "user",
+ "content": textwrap.dedent(f"""\
+ I am building a fine-tuning dataset for an LLM. I need you to write a \
+ system prompt that will be included in every training example to tell \
+ the model what task it is performing.
+
+ Here is the task information:
+ - Dataset type: {dtype}
+ - Task description: {task_desc}
+ - The USER (input) columns are: {user_cols}
+ - The ASSISTANT (output) columns are: {asst_cols}
+ {label_info}
+
+ Write a system prompt that:
+ 1. Explains what task the model is performing in plain language
+ 2. Describes what input it will receive
+ 3. Describes what output it should produce
+ 4. Is 2-4 sentences long
+
+ Write ONLY the system prompt text. No quotes, no labels, no explanation around it."""),
+ },
+ ]
+ raw3 = _generate_with_backend(backend, messages3, max_tokens=256)
+ print(f"🤖 Pass 3 done ({time.monotonic() - t3:.1f}s): {raw3[:200] if raw3 else None}", flush=True)
+
+ if raw3:
+ # Pass 3 returns raw text, not JSON — clean it up
+ cleaned = raw3.strip().strip('"').strip("'").strip()
+ if len(cleaned) >= 20 and cleaned.lower() not in ("null", "none", ""):
+ sys_prompt = cleaned
+
+ # Build suggested_mapping (column → role, for the frontend dropdowns)
+ suggested_mapping = {}
+ for col, role in column_roles.items():
+ if col in columns and role in ("user", "assistant", "system"):
+ suggested_mapping[col] = role
+
+ # Build user notification from Pass 1 classification
+ desc = pass1.get("task_description") or pass1.get("description", "")
+ note_parts = [f"This is a {dtype} dataset (not conversational)."]
+ if desc:
+ note_parts.append(desc)
+ note_parts.append("Columns have been mapped to conversation roles. You can adjust the mapping if needed.")
+ user_notification = " ".join(note_parts)
+
+ total_time = time.monotonic() - t0
+ print(
+ f"🤖 Advisor complete ({total_time:.1f}s): type={dtype}, "
+ f"mapping={suggested_mapping}, sys_prompt={bool(sys_prompt)}, label_map={bool(label_map)}",
+ flush=True,
+ )
+
+ return {
+ "success": True,
+ "suggested_mapping": suggested_mapping,
+ "system_prompt": sys_prompt,
+ "label_mapping": label_map if label_map else None,
+ "dataset_type": dtype,
+ "is_conversational": is_conv,
+ "user_notification": user_notification,
+ }
+
+ except Exception as e:
+ logger.warning(f"Advisor multi-pass failed: {e}")
+ return None
+
+ finally:
+ if backend is not None:
+ try:
+ backend.unload_model()
+ print("🤖 Advisor model unloaded")
+ except Exception:
+ pass
+
+
+def llm_conversion_advisor(
+ column_names: list[str],
+ 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.
+
+ Falls back to simple llm_classify_columns() if the multi-pass advisor fails.
+
+ Returns:
+ Dict with keys: success, suggested_mapping, system_prompt, user_template,
+ assistant_template, label_mapping, dataset_type, is_conversational,
+ user_notification. Or None on complete failure.
+ """
+ # Fetch HF dataset card if this looks like a HF dataset (has a slash)
+ dataset_card = None
+ dataset_metadata = None
+ if dataset_name and "/" in dataset_name:
+ dataset_card, dataset_metadata = fetch_hf_dataset_card(dataset_name, hf_token)
+
+ # Try multi-pass advisor
+ result = _run_multi_pass_advisor(
+ columns=column_names,
+ samples=samples,
+ 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"):
+ print(f"🤖 Conversion advisor succeeded: type={result.get('dataset_type')}")
+ return result
+
+ # Fallback: simple column classification
+ logger.info("Advisor failed, falling back to simple column classification")
+ simple_mapping = llm_classify_columns(column_names, samples)
+ if simple_mapping:
+ return {
+ "success": True,
+ "suggested_mapping": {
+ col: role for col, role in simple_mapping.items()
+ if role in ("user", "assistant", "system")
+ },
+ "dataset_type": None,
+ "is_conversational": None,
+ "user_notification": None,
+ }
+
+ return None
diff --git a/studio/backend/utils/datasets/vlm_processing.py b/studio/backend/utils/datasets/vlm_processing.py
index b17c57ef35..4caf80cbbb 100644
--- a/studio/backend/utils/datasets/vlm_processing.py
+++ b/studio/backend/utils/datasets/vlm_processing.py
@@ -9,6 +9,7 @@ for VLM datasets based on content analysis and heuristics.
"""
import re
+from itertools import islice
def generate_smart_vlm_instruction(
@@ -176,7 +177,46 @@ def generate_smart_vlm_instruction(
"confidence": 0.75,
}
- # ===== LEVEL 4: Generic Fallback =====
+ # ===== LEVEL 4: LLM-Assisted Instruction Generation =====
+ try:
+ from .llm_assist import llm_generate_vlm_instruction
+
+ sample_rows = []
+ for s in islice(dataset, 5):
+ row = {}
+ for col in s:
+ val = s[col]
+ if hasattr(val, 'size') and hasattr(val, 'mode'): # PIL Image
+ row[col] = ""
+ elif isinstance(val, list):
+ row[col] = str(val)[:300]
+ else:
+ row[col] = str(val)[:300]
+ sample_rows.append(row)
+
+ llm_result = llm_generate_vlm_instruction(
+ column_names=list(column_names),
+ samples=sample_rows,
+ dataset_name=dataset_name,
+ )
+ if llm_result and llm_result.get("instruction"):
+ print(
+ f"\n[DEBUG] LLM-assisted VLM instruction generated: "
+ f"'{llm_result['instruction']}' (confidence={llm_result.get('confidence', 'N/A')})\n",
+ flush=True,
+ )
+ return {
+ "instruction": llm_result["instruction"],
+ "instruction_column": None,
+ "instruction_type": "llm_assisted",
+ "uses_dynamic_instruction": False,
+ "confidence": llm_result.get("confidence", 0.85),
+ }
+ except Exception as e:
+ import logging
+ logging.getLogger(__name__).debug(f"LLM-assisted instruction skipped: {e}")
+
+ # ===== LEVEL 5: Generic Fallback =====
return {
"instruction": "Describe this image in detail.",
"instruction_column": None,
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..e4079663a8 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,11 @@ export function DatasetMappingCard({
isVlm = false,
isAudio = false,
format,
+ onAiAssist,
+ isAiLoading = false,
+ aiError,
+ advisorNotification,
+ advisorSystemPrompt,
}: {
mapping: Record;
mappingOk: boolean;
@@ -103,6 +109,11 @@ export function DatasetMappingCard({
isVlm?: boolean;
isAudio?: boolean;
format?: string;
+ onAiAssist?: () => void;
+ isAiLoading?: boolean;
+ aiError?: string | null;
+ advisorNotification?: string | null;
+ advisorSystemPrompt?: string;
}) {
const entries = Object.entries(mapping);
const requiredLabel = isAudio
@@ -157,7 +168,7 @@ export function DatasetMappingCard({
>
{mappingOk
? autoDetected
- ? "We auto-detected the column mapping below. You can change it using the dropdowns in the column headers."
+ ? "We auto-detected the column mapping below. You can change it using the dropdowns in the column headers, or use AI Assist for a smarter mapping strategy."
: "Looks good. We'll convert this dataset automatically."
: `Assign roles to columns using the dropdowns in the headers. At minimum, assign ${requiredLabel}.`}
@@ -181,6 +192,47 @@ export function DatasetMappingCard({
Use the dropdowns in the column headers to assign roles.
)}
+ {onAiAssist && (
+
+
+ {aiError && (
+
{aiError}
+ )}
+
+ )}
+ {advisorNotification && (
+
+
+
+ {advisorNotification}
+
+ {advisorSystemPrompt && (
+
+ System:{" "}
+ {advisorSystemPrompt}
+
+ )}
+
+ )}
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 650d03cc14..c79cf77e90 100644
--- a/studio/frontend/src/features/studio/sections/dataset-preview-dialog.tsx
+++ b/studio/frontend/src/features/studio/sections/dataset-preview-dialog.tsx
@@ -2,6 +2,7 @@
// Copyright © 2025 Unsloth AI
import { type ReactNode, useCallback, useEffect, useMemo, useRef, useState } from "react";
+import { aiAssistMapping } from "@/features/training/api/datasets-api";
import type { ColumnDef } from "@tanstack/react-table";
import {
Dialog,
@@ -29,6 +30,12 @@ import {
remapRolesForFormat,
} from "./dataset-preview-dialog-mapping";
+/** Chatml → format-specific role remap (only for formats that differ from chatml). */
+const ROLE_REMAP: Record> = {
+ alpaca: { user: "instruction", system: "input", assistant: "output" },
+ sharegpt: { user: "human", assistant: "gpt", system: "system" },
+};
+
type DatasetPreviewDialogProps = {
open: boolean;
onOpenChange: (open: boolean) => void;
@@ -58,11 +65,22 @@ export function DatasetPreviewDialog({
const [loading, setLoading] = useState(false);
const [error, setError] = useState(null);
- const { manualMapping, setManualMapping, datasetFormat } = useTrainingConfigStore(
+ const {
+ manualMapping, setManualMapping, datasetFormat,
+ setDatasetAdvisorFields, datasetAdvisorNotification,
+ datasetSystemPrompt,
+ selectedModel,
+ modelType,
+ } = useTrainingConfigStore(
useShallow((s) => ({
manualMapping: s.datasetManualMapping,
setManualMapping: s.setDatasetManualMapping,
datasetFormat: s.datasetFormat,
+ setDatasetAdvisorFields: s.setDatasetAdvisorFields,
+ datasetAdvisorNotification: s.datasetAdvisorNotification,
+ datasetSystemPrompt: s.datasetSystemPrompt,
+ selectedModel: s.selectedModel,
+ modelType: s.modelType,
})),
);
const { isStarting, startError, startTrainingRun } = useTrainingActions();
@@ -79,6 +97,52 @@ export function DatasetPreviewDialog({
const availableRoles = getAvailableRoles(effectiveIsVlm, datasetFormat, effectiveIsAudio);
const isHfDataset = datasetSource === "huggingface";
+ // ── AI Assist ──────────────────────────────────────────────────────
+ const [isAiLoading, setIsAiLoading] = useState(false);
+ const [aiError, setAiError] = useState(null);
+
+ const handleAiAssist = useCallback(async () => {
+ if (!data?.columns || !data?.preview_samples) return;
+ setIsAiLoading(true);
+ setAiError(null);
+
+ try {
+ const result = await aiAssistMapping({
+ columns: data.columns,
+ samples: data.preview_samples,
+ datasetName: datasetName,
+ hfToken: hfToken,
+ modelName: selectedModel,
+ modelType: modelType,
+ });
+
+ if (result.success && result.suggested_mapping) {
+ // Remap from chatml roles (user/assistant/system) to format-specific roles
+ const table = ROLE_REMAP[datasetFormat];
+ const mapped: Record = {};
+ for (const [col, role] of Object.entries(result.suggested_mapping)) {
+ mapped[col] = table ? (table[role] ?? role) : role;
+ }
+ setManualMapping(mapped);
+
+ // Store conversion advisor fields (system prompt, label mapping, notification)
+ if (result.system_prompt || result.label_mapping || result.user_notification) {
+ setDatasetAdvisorFields({
+ systemPrompt: result.system_prompt ?? undefined,
+ labelMapping: result.label_mapping ?? undefined,
+ notification: result.user_notification ?? null,
+ });
+ }
+ } else {
+ setAiError(result.warning || "AI could not determine column roles.");
+ }
+ } catch (err) {
+ setAiError(err instanceof Error ? err.message : "AI assist failed.");
+ } finally {
+ setIsAiLoading(false);
+ }
+ }, [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);
useEffect(() => {
@@ -180,7 +244,8 @@ export function DatasetPreviewDialog({
// Build TanStack Table columns from the column names
const tableColumns = useMemo>[]>(() => {
if (!columns.length) return [];
- return columns.map((colName) => ({
+
+ const dataCols: ColumnDef>[] = columns.map((colName) => ({
accessorKey: colName,
header: () => (
@@ -247,12 +312,42 @@ export function DatasetPreviewDialog({
);
},
}));
+
+ // Prepend generated system prompt column when advisor is active
+ if (datasetSystemPrompt) {
+ dataCols.unshift({
+ id: "__system_generated",
+ header: () => (
+
+
+ System (generated)
+
+ {mappingEnabled && (
+
+ System
+
+ )}
+
+ ),
+ cell: () => (
+
+ {datasetSystemPrompt}
+
+ ),
+ });
+ }
+
+ return dataCols;
}, [
columns,
manualMapping,
handleRoleChange,
mappingEnabled,
availableRoles,
+ datasetSystemPrompt,
]);
return (
@@ -274,7 +369,7 @@ export function DatasetPreviewDialog({
{/* Body */}
-
+
{/* Loading */}
{loading && (
@@ -361,11 +456,16 @@ export function DatasetPreviewDialog({
isVlm={effectiveIsVlm}
isAudio={effectiveIsAudio}
format={datasetFormat}
+ onAiAssist={handleAiAssist}
+ isAiLoading={isAiLoading}
+ aiError={aiError}
+ advisorNotification={datasetAdvisorNotification}
+ advisorSystemPrompt={datasetSystemPrompt || undefined}
/>
)}
{/* Data table */}
-
+
diff --git a/studio/frontend/src/features/training/api/datasets-api.ts b/studio/frontend/src/features/training/api/datasets-api.ts
index 047925e60a..3ff04189f7 100644
--- a/studio/frontend/src/features/training/api/datasets-api.ts
+++ b/studio/frontend/src/features/training/api/datasets-api.ts
@@ -62,6 +62,58 @@ export async function uploadTrainingDataset(
return res.json();
}
+// ── AI Assist ────────────────────────────────────────────────────────
+
+type AiAssistMappingArgs = {
+ columns: string[];
+ samples: Record
[];
+ datasetName?: string | null;
+ hfToken?: string | null;
+ modelName?: string | null;
+ modelType?: "text" | "vision" | "audio" | "embeddings" | null;
+};
+
+export type AiAssistMappingResponse = {
+ success: boolean;
+ suggested_mapping?: Record | null;
+ warning?: string | null;
+ // Conversion advisor fields
+ system_prompt?: string | null;
+ label_mapping?: Record> | null;
+ dataset_type?: string | null;
+ is_conversational?: boolean | null;
+ user_notification?: string | null;
+};
+
+export async function aiAssistMapping({
+ columns,
+ samples,
+ datasetName,
+ hfToken,
+ modelName,
+ modelType,
+}: AiAssistMappingArgs): Promise {
+ const res = await authFetch("/api/datasets/ai-assist-mapping", {
+ method: "POST",
+ headers: { "Content-Type": "application/json" },
+ body: JSON.stringify({
+ columns,
+ samples: samples.slice(0, 5),
+ dataset_name: datasetName || undefined,
+ hf_token: hfToken || undefined,
+ model_name: modelName || undefined,
+ model_type: modelType || undefined,
+ }),
+ });
+
+ if (!res.ok) {
+ const body = await res.json().catch(() => null);
+ throw new Error(body?.detail || `AI assist failed (${res.status})`);
+ }
+
+ return res.json();
+}
+
export async function listLocalDatasets(): Promise {
const res = await authFetch("/api/datasets/local");
if (!res.ok) {
diff --git a/studio/frontend/src/features/training/api/mappers.ts b/studio/frontend/src/features/training/api/mappers.ts
index b10c282113..38ec92169e 100644
--- a/studio/frontend/src/features/training/api/mappers.ts
+++ b/studio/frontend/src/features/training/api/mappers.ts
@@ -31,8 +31,23 @@ export function buildTrainingStartPayload(
config.datasetSource === "upload" && config.uploadedFile
? [config.uploadedFile]
: [];
- const customFormatMapping =
- Object.keys(config.datasetManualMapping).length > 0 ? config.datasetManualMapping : undefined;
+ let customFormatMapping: Record | undefined =
+ Object.keys(config.datasetManualMapping).length > 0
+ ? { ...config.datasetManualMapping }
+ : undefined;
+
+ // Inject conversion advisor metadata into the mapping (__ prefix keys)
+ const hasAdvisorMeta =
+ config.datasetSystemPrompt ||
+ Object.keys(config.datasetLabelMapping).length > 0;
+ if (customFormatMapping && hasAdvisorMeta) {
+ if (config.datasetSystemPrompt) {
+ customFormatMapping.__system_prompt = config.datasetSystemPrompt;
+ }
+ if (Object.keys(config.datasetLabelMapping).length > 0) {
+ customFormatMapping.__label_mapping = config.datasetLabelMapping;
+ }
+ }
return {
model_name: config.selectedModel ?? "",
diff --git a/studio/frontend/src/features/training/stores/training-config-store.ts b/studio/frontend/src/features/training/stores/training-config-store.ts
index cbaf9f375c..9d379edb7d 100644
--- a/studio/frontend/src/features/training/stores/training-config-store.ts
+++ b/studio/frontend/src/features/training/stores/training-config-store.ts
@@ -31,6 +31,11 @@ const initialState: TrainingConfigState = {
datasetSplit: null,
datasetEvalSplit: null,
datasetManualMapping: emptyManualMapping(),
+ datasetSystemPrompt: "",
+ datasetUserTemplate: "",
+ datasetAssistantTemplate: "",
+ datasetLabelMapping: {},
+ datasetAdvisorNotification: null,
datasetSliceStart: null,
datasetSliceEnd: null,
uploadedFile: null,
@@ -241,6 +246,11 @@ export const useTrainingConfigStore = create()(
datasetSplit: null,
datasetEvalSplit: null,
datasetManualMapping: emptyManualMapping(),
+ datasetSystemPrompt: "",
+ datasetUserTemplate: "",
+ datasetAssistantTemplate: "",
+ datasetLabelMapping: {},
+ datasetAdvisorNotification: null,
datasetSliceStart: null,
datasetSliceEnd: null,
isDatasetImage: null,
@@ -404,6 +414,22 @@ export const useTrainingConfigStore = create()(
},
setDatasetManualMapping: (datasetManualMapping) =>
set({ datasetManualMapping }),
+ setDatasetAdvisorFields: (fields) =>
+ set({
+ datasetSystemPrompt: fields.systemPrompt ?? get().datasetSystemPrompt,
+ datasetUserTemplate: "", // templates no longer used
+ datasetAssistantTemplate: "", // templates no longer used
+ datasetLabelMapping: fields.labelMapping ?? get().datasetLabelMapping,
+ datasetAdvisorNotification: fields.notification !== undefined ? fields.notification : get().datasetAdvisorNotification,
+ }),
+ clearDatasetAdvisorFields: () =>
+ set({
+ datasetSystemPrompt: "",
+ datasetUserTemplate: "",
+ datasetAssistantTemplate: "",
+ datasetLabelMapping: {},
+ datasetAdvisorNotification: null,
+ }),
setDatasetSliceStart: (datasetSliceStart) => set({ datasetSliceStart }),
setDatasetSliceEnd: (datasetSliceEnd) => set({ datasetSliceEnd }),
setUploadedFile: (uploadedFile) => {
@@ -479,7 +505,7 @@ export const useTrainingConfigStore = create()(
},
{
name: "unsloth_training_config_v1",
- version: 7,
+ version: 8,
migrate: (persisted, version) => {
const s = persisted as Record;
if (version < 2 && s.datasetSubset == null && s.datasetConfig != null) {
@@ -502,6 +528,13 @@ export const useTrainingConfigStore = create()(
s.datasetSliceStart ??= null;
s.datasetSliceEnd ??= null;
}
+ if (version < 8) {
+ s.datasetSystemPrompt ??= "";
+ s.datasetUserTemplate ??= "";
+ s.datasetAssistantTemplate ??= "";
+ s.datasetLabelMapping ??= {};
+ s.datasetAdvisorNotification ??= null;
+ }
return s as unknown as TrainingConfigStore;
},
partialize: partializePersistedState,
diff --git a/studio/frontend/src/features/training/types/api.ts b/studio/frontend/src/features/training/types/api.ts
index 54d5b98e72..16de48de10 100644
--- a/studio/frontend/src/features/training/types/api.ts
+++ b/studio/frontend/src/features/training/types/api.ts
@@ -17,7 +17,7 @@ export interface TrainingStartRequest {
dataset_slice_end: number | null;
local_datasets: string[];
format_type: string;
- custom_format_mapping?: Record | null;
+ custom_format_mapping?: Record | null;
num_epochs: number;
learning_rate: string;
batch_size: number;
diff --git a/studio/frontend/src/features/training/types/config.ts b/studio/frontend/src/features/training/types/config.ts
index 32069baa68..cd4e188002 100644
--- a/studio/frontend/src/features/training/types/config.ts
+++ b/studio/frontend/src/features/training/types/config.ts
@@ -29,6 +29,11 @@ export interface TrainingConfigState {
datasetSplit: string | null;
datasetEvalSplit: string | null;
datasetManualMapping: DatasetManualMapping;
+ datasetSystemPrompt: string;
+ datasetUserTemplate: string;
+ datasetAssistantTemplate: string;
+ datasetLabelMapping: Record>;
+ datasetAdvisorNotification: string | null;
datasetSliceStart: string | null;
datasetSliceEnd: string | null;
uploadedFile: string | null;
@@ -95,6 +100,12 @@ export interface TrainingConfigActions {
setDatasetSplit: (split: string | null) => void;
setDatasetEvalSplit: (split: string | null) => void;
setDatasetManualMapping: (mapping: DatasetManualMapping) => void;
+ setDatasetAdvisorFields: (fields: {
+ systemPrompt?: string;
+ labelMapping?: Record>;
+ notification?: string | null;
+ }) => void;
+ clearDatasetAdvisorFields: () => void;
setDatasetSliceStart: (value: string | null) => void;
setDatasetSliceEnd: (value: string | null) => void;
setUploadedFile: (file: string | null) => void;