variable changes and some cleanup
This commit is contained in:
parent
87f2b2a9db
commit
f04c684d8a
19 changed files with 93 additions and 111 deletions
|
|
@ -30,11 +30,6 @@ from trl import SFTTrainer, SFTConfig
|
|||
logging.basicConfig(level=logging.INFO)
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Process-level flag: set True after CUDA-heavy audio preprocessing (Whisper/DAC/BiCodec).
|
||||
# Once CUDA has been used for audio processing, fork-based multiprocessing (num_proc>1)
|
||||
# deadlocks because forked children inherit CUDA's internal thread locks.
|
||||
# This flag is never reset — once contaminated, the process stays contaminated.
|
||||
_CUDA_AUDIO_PREPROCESSING_DONE = False
|
||||
|
||||
@dataclass
|
||||
class TrainingProgress:
|
||||
|
|
@ -75,6 +70,7 @@ class UnslothTrainer:
|
|||
self.is_audio = False
|
||||
self.is_audio_vlm = False # Multimodal model (e.g. Gemma 3N) trained on audio data
|
||||
self._audio_type = None # 'csm', 'whisper', 'snac', 'xcodec2', 'bicodec', 'dac'
|
||||
self._cuda_audio_used = False # Set once after audio CUDA preprocessing; never cleared
|
||||
self._spark_tts_repo_dir = None # Path to downloaded Spark-TTS repo (for BiCodecTokenizer)
|
||||
self.model_name = None
|
||||
|
||||
|
|
@ -335,7 +331,7 @@ class UnslothTrainer:
|
|||
max_seq_length: int = 2048,
|
||||
load_in_4bit: bool = True,
|
||||
hf_token: Optional[str] = None,
|
||||
is_dataset_multimodal: bool = False,
|
||||
is_dataset_image: bool = False,
|
||||
is_dataset_audio: bool = False) -> bool:
|
||||
"""Load model for training (supports both text and vision models)"""
|
||||
try:
|
||||
|
|
@ -381,14 +377,14 @@ class UnslothTrainer:
|
|||
# Uses FastModel + SFTTrainer with audio collator
|
||||
self.is_audio_vlm = not self.is_audio and is_vision_model(model_name) and is_dataset_audio
|
||||
# VLM: vision model with image dataset (mutually exclusive with audio VLM)
|
||||
self.is_vlm = not self.is_audio and not self.is_audio_vlm and is_vision_model(model_name) and is_dataset_multimodal
|
||||
self.is_vlm = not self.is_audio and not self.is_audio_vlm and is_vision_model(model_name) and is_dataset_image
|
||||
self.model_name = model_name
|
||||
self.max_seq_length = max_seq_length
|
||||
|
||||
logger.info(f"Audio type: {self._audio_type}")
|
||||
if not self.is_audio:
|
||||
logger.info(f"Model architecture is vision: {is_vision_model(model_name)}")
|
||||
logger.info(f"Dataset is multimodal: {is_dataset_multimodal}, audio: {is_dataset_audio}")
|
||||
logger.info(f"Dataset has images: {is_dataset_image}, audio: {is_dataset_audio}")
|
||||
logger.info(f"Using VLM path: {self.is_vlm}, Audio VLM: {self.is_audio_vlm}")
|
||||
|
||||
# Reset training state for new run
|
||||
|
|
@ -555,10 +551,26 @@ class UnslothTrainer:
|
|||
print("Model loaded successfully")
|
||||
return True
|
||||
|
||||
except OSError as e:
|
||||
if "could not get source code" in str(e) and not getattr(self, '_source_code_retried', False):
|
||||
# Unsloth's patching can leave stale state that makes
|
||||
# inspect.getsource() fail when switching model families
|
||||
# (e.g. gemma3 → gemma3n). The load always succeeds on a
|
||||
# second attempt because the failed first call's partial
|
||||
# imports clean up the stale state as a side effect.
|
||||
self._source_code_retried = True
|
||||
print(f"\n'could not get source code' — retrying once...\n")
|
||||
return self.load_model(model_name, max_seq_length, load_in_4bit, hf_token,
|
||||
is_dataset_image, is_dataset_audio)
|
||||
logger.error(f"Error loading model: {e}")
|
||||
self._update_progress(error=str(e), is_training=False)
|
||||
return False
|
||||
except Exception as e:
|
||||
logger.error(f"Error loading model: {e}")
|
||||
self._update_progress(error=str(e), is_training=False)
|
||||
return False
|
||||
finally:
|
||||
self._source_code_retried = False
|
||||
|
||||
def prepare_model_for_training(self,
|
||||
use_lora: bool = True,
|
||||
|
|
@ -1204,9 +1216,7 @@ class UnslothTrainer:
|
|||
import gc
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
global _CUDA_AUDIO_PREPROCESSING_DONE
|
||||
_CUDA_AUDIO_PREPROCESSING_DONE = True
|
||||
self._cuda_audio_used = True
|
||||
|
||||
if not processed_examples:
|
||||
raise ValueError(
|
||||
|
|
@ -1394,9 +1404,7 @@ class UnslothTrainer:
|
|||
import gc
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
global _CUDA_AUDIO_PREPROCESSING_DONE
|
||||
_CUDA_AUDIO_PREPROCESSING_DONE = True
|
||||
self._cuda_audio_used = True
|
||||
|
||||
if not processed_examples:
|
||||
raise ValueError(
|
||||
|
|
@ -1579,13 +1587,7 @@ class UnslothTrainer:
|
|||
import gc
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
# Mark process as CUDA-contaminated from audio preprocessing.
|
||||
# Fork-based multiprocessing (num_proc>1) will deadlock after this
|
||||
# because forked children inherit CUDA's internal thread locks from
|
||||
# Whisper/DAC processing that can't be released.
|
||||
global _CUDA_AUDIO_PREPROCESSING_DONE
|
||||
_CUDA_AUDIO_PREPROCESSING_DONE = True
|
||||
self._cuda_audio_used = True
|
||||
|
||||
if not processed_examples:
|
||||
raise ValueError(
|
||||
|
|
@ -2236,7 +2238,7 @@ class UnslothTrainer:
|
|||
"output_dir": output_dir,
|
||||
"report_to": ["wandb"] if training_args.get('enable_wandb', False) else "none",
|
||||
"include_num_input_tokens_seen": True, # Enable token counting
|
||||
"dataset_num_proc": safe_num_proc(max(1, os.cpu_count() // 4)),
|
||||
"dataset_num_proc": 1 if (self.is_audio or self.is_audio_vlm or self._cuda_audio_used) else safe_num_proc(max(1, os.cpu_count() // 4)),
|
||||
"max_seq_length": training_args.get('max_seq_length', 2048),
|
||||
}
|
||||
|
||||
|
|
@ -2424,19 +2426,11 @@ class UnslothTrainer:
|
|||
try:
|
||||
from unsloth.chat_templates import train_on_responses_only
|
||||
|
||||
# After CUDA-heavy audio preprocessing (Whisper/DAC/BiCodec/SNAC),
|
||||
# fork-based multiprocessing deadlocks because children inherit
|
||||
# CUDA's internal thread locks. Use single-process mode instead.
|
||||
toro_num_proc = config_args.get("dataset_num_proc", safe_num_proc(max(1, os.cpu_count() // 4)))
|
||||
if _CUDA_AUDIO_PREPROCESSING_DONE:
|
||||
toro_num_proc = 1
|
||||
print("Using single-process train_on_responses (CUDA audio preprocessing detected)\n")
|
||||
|
||||
self.trainer = train_on_responses_only(
|
||||
self.trainer,
|
||||
instruction_part=instruction_part,
|
||||
response_part=response_part,
|
||||
num_proc=toro_num_proc,
|
||||
num_proc=config_args["dataset_num_proc"],
|
||||
)
|
||||
print("Train on responses only configured successfully\n")
|
||||
|
||||
|
|
|
|||
|
|
@ -116,7 +116,7 @@ class TrainingBackend:
|
|||
train_split: str = "train",
|
||||
eval_split: str = None,
|
||||
eval_steps: float = 0.00,
|
||||
is_dataset_multimodal: bool = False,
|
||||
is_dataset_image: bool = False,
|
||||
is_dataset_audio: bool = False) -> bool:
|
||||
"""
|
||||
Start training.
|
||||
|
|
@ -146,9 +146,6 @@ class TrainingBackend:
|
|||
import torch as _torch
|
||||
if _torch.cuda.is_available():
|
||||
_torch.cuda.synchronize()
|
||||
# Reset torch dynamo/compiler caches — Unsloth's compiled SFTTrainer
|
||||
# and model.for_training() set class-level state that persists between
|
||||
# runs (e.g. BiCodec text trainer pollutes subsequent VLM runs).
|
||||
_torch._dynamo.reset()
|
||||
_torch.compiler.reset()
|
||||
import gc
|
||||
|
|
@ -182,7 +179,7 @@ class TrainingBackend:
|
|||
max_seq_length=max_seq_length,
|
||||
load_in_4bit=load_in_4bit if use_lora_actual else False, # Only 4bit for LoRA
|
||||
hf_token=hf_token if hf_token.strip() else None,
|
||||
is_dataset_multimodal=is_dataset_multimodal,
|
||||
is_dataset_image=is_dataset_image,
|
||||
is_dataset_audio=is_dataset_audio,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -27,7 +27,7 @@ class CheckFormatResponse(BaseModel):
|
|||
requires_manual_mapping: bool
|
||||
detected_format: str
|
||||
columns: List[str]
|
||||
is_multimodal: bool = False
|
||||
is_image: bool = False
|
||||
is_audio: bool = False
|
||||
multimodal_columns: Optional[List[str]] = None
|
||||
suggested_mapping: Optional[Dict[str, str]] = None
|
||||
|
|
|
|||
|
|
@ -65,7 +65,7 @@ class TrainingStartRequest(BaseModel):
|
|||
finetune_language_layers: bool = Field(False, description="Finetune language layers")
|
||||
finetune_attention_modules: bool = Field(False, description="Finetune attention modules")
|
||||
finetune_mlp_modules: bool = Field(False, description="Finetune MLP modules")
|
||||
is_dataset_multimodal: bool = Field(False, description="Whether the dataset contains multimodal (image) data")
|
||||
is_dataset_image: bool = Field(False, description="Whether the dataset contains image data")
|
||||
is_dataset_audio: bool = Field(False, description="Whether the dataset contains audio data")
|
||||
|
||||
# Logging parameters
|
||||
|
|
|
|||
|
|
@ -183,7 +183,7 @@ def check_format(request: CheckFormatRequest):
|
|||
# Run lightweight format check on the preview slice
|
||||
result = check_dataset_format(preview_slice, is_vlm=request.is_vlm)
|
||||
|
||||
logger.info(f"Format check result: requires_mapping={result['requires_manual_mapping']}, format={result['detected_format']}, is_multimodal={result.get('is_multimodal', False)}")
|
||||
logger.info(f"Format check result: requires_mapping={result['requires_manual_mapping']}, format={result['detected_format']}, is_image={result.get('is_image', False)}")
|
||||
|
||||
# Generate preview samples
|
||||
preview_samples = None
|
||||
|
|
@ -211,7 +211,7 @@ def check_format(request: CheckFormatRequest):
|
|||
requires_manual_mapping=result["requires_manual_mapping"],
|
||||
detected_format=result["detected_format"],
|
||||
columns=result["columns"],
|
||||
is_multimodal=result.get("is_multimodal", False),
|
||||
is_image=result.get("is_image", False),
|
||||
is_audio=result.get("is_audio", False),
|
||||
multimodal_columns=result.get("multimodal_columns"),
|
||||
suggested_mapping=result.get("suggested_mapping"),
|
||||
|
|
|
|||
|
|
@ -178,7 +178,7 @@ async def start_training(
|
|||
"finetune_language_layers": request.finetune_language_layers,
|
||||
"finetune_attention_modules": request.finetune_attention_modules,
|
||||
"finetune_mlp_modules": request.finetune_mlp_modules,
|
||||
"is_dataset_multimodal": request.is_dataset_multimodal,
|
||||
"is_dataset_image": request.is_dataset_image,
|
||||
"is_dataset_audio": request.is_dataset_audio,
|
||||
"enable_wandb": request.enable_wandb,
|
||||
"wandb_token": request.wandb_token or "",
|
||||
|
|
|
|||
|
|
@ -67,8 +67,8 @@ def check_dataset_format(dataset, is_vlm: bool = False) -> dict:
|
|||
multimodal_info = detect_multimodal_dataset(dataset)
|
||||
is_audio = multimodal_info.get("is_audio", False)
|
||||
|
||||
if multimodal_info["is_multimodal"] and not is_audio:
|
||||
is_vlm = True # Route to VLM detection for image datasets only
|
||||
if multimodal_info["is_image"]:
|
||||
is_vlm = True # Route to VLM detection for image datasets
|
||||
|
||||
# Common audio fields for all return paths
|
||||
audio_fields = {
|
||||
|
|
@ -88,7 +88,7 @@ def check_dataset_format(dataset, is_vlm: bool = False) -> dict:
|
|||
"suggested_mapping": None,
|
||||
"detected_image_column": vlm_structure.get("image_column"),
|
||||
"detected_text_column": vlm_structure.get("text_column"),
|
||||
"is_multimodal": multimodal_info["is_multimodal"],
|
||||
"is_image": multimodal_info["is_image"],
|
||||
"multimodal_columns": multimodal_info.get("multimodal_columns"),
|
||||
**audio_fields,
|
||||
}
|
||||
|
|
@ -105,7 +105,7 @@ def check_dataset_format(dataset, is_vlm: bool = False) -> dict:
|
|||
"suggested_mapping": None,
|
||||
"detected_image_column": None,
|
||||
"detected_text_column": multimodal_info.get("detected_text_column"),
|
||||
"is_multimodal": True,
|
||||
"is_image": False,
|
||||
"multimodal_columns": multimodal_info.get("audio_columns"),
|
||||
**audio_fields,
|
||||
}
|
||||
|
|
@ -124,7 +124,7 @@ def check_dataset_format(dataset, is_vlm: bool = False) -> dict:
|
|||
"suggested_mapping": heuristic_mapping,
|
||||
"detected_image_column": None,
|
||||
"detected_text_column": None,
|
||||
"is_multimodal": False,
|
||||
"is_image": False,
|
||||
"multimodal_columns": None,
|
||||
**audio_fields,
|
||||
}
|
||||
|
|
@ -136,7 +136,7 @@ def check_dataset_format(dataset, is_vlm: bool = False) -> dict:
|
|||
"suggested_mapping": None,
|
||||
"detected_image_column": None,
|
||||
"detected_text_column": None,
|
||||
"is_multimodal": False,
|
||||
"is_image": False,
|
||||
"multimodal_columns": None,
|
||||
**audio_fields,
|
||||
}
|
||||
|
|
@ -149,7 +149,7 @@ def check_dataset_format(dataset, is_vlm: bool = False) -> dict:
|
|||
"suggested_mapping": None,
|
||||
"detected_image_column": None,
|
||||
"detected_text_column": None,
|
||||
"is_multimodal": False,
|
||||
"is_image": False,
|
||||
"multimodal_columns": None,
|
||||
**audio_fields,
|
||||
}
|
||||
|
|
@ -278,7 +278,7 @@ def format_dataset(
|
|||
"chat_column": chat_column,
|
||||
"is_standardized": True,
|
||||
"requires_manual_mapping": False,
|
||||
"is_multimodal": multimodal_info["is_multimodal"],
|
||||
"is_image": multimodal_info["is_image"],
|
||||
"multimodal_info": multimodal_info,
|
||||
"warnings": [f"Applied user-provided column mapping ({format_type}): {custom_format_mapping}"]
|
||||
}
|
||||
|
|
@ -290,7 +290,7 @@ def format_dataset(
|
|||
"chat_column": None,
|
||||
"is_standardized": False,
|
||||
"requires_manual_mapping": True,
|
||||
"is_multimodal": multimodal_info["is_multimodal"],
|
||||
"is_image": multimodal_info["is_image"],
|
||||
"multimodal_info": multimodal_info,
|
||||
"warnings": [f"Failed to apply user mapping: {e}"]
|
||||
}
|
||||
|
|
@ -301,7 +301,7 @@ def format_dataset(
|
|||
warnings = []
|
||||
|
||||
# Add multimodal warning if detected
|
||||
if multimodal_info["is_multimodal"]:
|
||||
if multimodal_info["is_image"]:
|
||||
warnings.append(
|
||||
f"Multimodal dataset detected. Found columns: {multimodal_info['multimodal_columns']}"
|
||||
)
|
||||
|
|
@ -318,7 +318,7 @@ def format_dataset(
|
|||
"chat_column": None,
|
||||
"is_standardized": True,
|
||||
"requires_manual_mapping": False,
|
||||
"is_multimodal": multimodal_info["is_multimodal"],
|
||||
"is_image": multimodal_info["is_image"],
|
||||
"multimodal_info": multimodal_info,
|
||||
"warnings": []
|
||||
}
|
||||
|
|
@ -338,7 +338,7 @@ def format_dataset(
|
|||
"chat_column": detected["chat_column"],
|
||||
"is_standardized": True,
|
||||
"requires_manual_mapping": False,
|
||||
"is_multimodal": multimodal_info["is_multimodal"],
|
||||
"is_image": multimodal_info["is_image"],
|
||||
"multimodal_info": multimodal_info,
|
||||
"warnings": []
|
||||
}
|
||||
|
|
@ -351,7 +351,7 @@ def format_dataset(
|
|||
"chat_column": detected["chat_column"],
|
||||
"is_standardized": False,
|
||||
"requires_manual_mapping": True,
|
||||
"is_multimodal": multimodal_info["is_multimodal"],
|
||||
"is_image": multimodal_info["is_image"],
|
||||
"multimodal_info": multimodal_info,
|
||||
"warnings": warnings
|
||||
}
|
||||
|
|
@ -364,7 +364,7 @@ def format_dataset(
|
|||
"chat_column": detected["chat_column"],
|
||||
"is_standardized": True,
|
||||
"requires_manual_mapping": False,
|
||||
"is_multimodal": multimodal_info["is_multimodal"],
|
||||
"is_image": multimodal_info["is_image"],
|
||||
"multimodal_info": multimodal_info,
|
||||
"warnings": warnings
|
||||
}
|
||||
|
|
@ -415,7 +415,7 @@ def format_dataset(
|
|||
"chat_column": "conversations",
|
||||
"is_standardized": True,
|
||||
"requires_manual_mapping": False,
|
||||
"is_multimodal": multimodal_info["is_multimodal"],
|
||||
"is_image": multimodal_info["is_image"],
|
||||
"multimodal_info": multimodal_info,
|
||||
"warnings": warnings
|
||||
}
|
||||
|
|
@ -438,7 +438,7 @@ def format_dataset(
|
|||
"chat_column": detected["chat_column"],
|
||||
"is_standardized": True,
|
||||
"requires_manual_mapping": False,
|
||||
"is_multimodal": multimodal_info["is_multimodal"],
|
||||
"is_image": multimodal_info["is_image"],
|
||||
"multimodal_info": multimodal_info,
|
||||
"warnings": warnings
|
||||
}
|
||||
|
|
@ -453,7 +453,7 @@ def format_dataset(
|
|||
"chat_column": detected["chat_column"],
|
||||
"is_standardized": False,
|
||||
"requires_manual_mapping": True,
|
||||
"is_multimodal": multimodal_info["is_multimodal"],
|
||||
"is_image": multimodal_info["is_image"],
|
||||
"multimodal_info": multimodal_info,
|
||||
"warnings": warnings
|
||||
}
|
||||
|
|
@ -469,7 +469,7 @@ def format_dataset(
|
|||
"chat_column": None,
|
||||
"is_standardized": True,
|
||||
"requires_manual_mapping": False,
|
||||
"is_multimodal": multimodal_info["is_multimodal"],
|
||||
"is_image": multimodal_info["is_image"],
|
||||
"multimodal_info": multimodal_info,
|
||||
"warnings": []
|
||||
}
|
||||
|
|
@ -492,7 +492,7 @@ def format_dataset(
|
|||
"chat_column": None,
|
||||
"is_standardized": True,
|
||||
"requires_manual_mapping": False,
|
||||
"is_multimodal": multimodal_info["is_multimodal"],
|
||||
"is_image": multimodal_info["is_image"],
|
||||
"multimodal_info": multimodal_info,
|
||||
"warnings": []
|
||||
}
|
||||
|
|
@ -506,7 +506,7 @@ def format_dataset(
|
|||
"chat_column": detected["chat_column"],
|
||||
"is_standardized": False,
|
||||
"requires_manual_mapping": True,
|
||||
"is_multimodal": multimodal_info["is_multimodal"],
|
||||
"is_image": multimodal_info["is_image"],
|
||||
"multimodal_info": multimodal_info,
|
||||
"warnings": warnings
|
||||
}
|
||||
|
|
@ -523,7 +523,7 @@ def format_dataset(
|
|||
"chat_column": "conversations",
|
||||
"is_standardized": True,
|
||||
"requires_manual_mapping": False,
|
||||
"is_multimodal": multimodal_info["is_multimodal"],
|
||||
"is_image": multimodal_info["is_image"],
|
||||
"multimodal_info": multimodal_info,
|
||||
"warnings": []
|
||||
}
|
||||
|
|
@ -541,7 +541,7 @@ def format_dataset(
|
|||
"chat_column": detected["chat_column"],
|
||||
"is_standardized": True,
|
||||
"requires_manual_mapping": False,
|
||||
"is_multimodal": multimodal_info["is_multimodal"],
|
||||
"is_image": multimodal_info["is_image"],
|
||||
"multimodal_info": multimodal_info,
|
||||
"warnings": []
|
||||
}
|
||||
|
|
@ -554,7 +554,7 @@ def format_dataset(
|
|||
"chat_column": detected["chat_column"],
|
||||
"is_standardized": True,
|
||||
"requires_manual_mapping": False,
|
||||
"is_multimodal": multimodal_info["is_multimodal"],
|
||||
"is_image": multimodal_info["is_image"],
|
||||
"multimodal_info": multimodal_info,
|
||||
"warnings": []
|
||||
}
|
||||
|
|
@ -575,7 +575,7 @@ def format_dataset(
|
|||
"chat_column": detected["chat_column"],
|
||||
"is_standardized": True,
|
||||
"requires_manual_mapping": False,
|
||||
"is_multimodal": multimodal_info["is_multimodal"],
|
||||
"is_image": multimodal_info["is_image"],
|
||||
"multimodal_info": multimodal_info,
|
||||
"warnings": warnings
|
||||
}
|
||||
|
|
@ -589,7 +589,7 @@ def format_dataset(
|
|||
"chat_column": detected["chat_column"],
|
||||
"is_standardized": False,
|
||||
"requires_manual_mapping": True,
|
||||
"is_multimodal": multimodal_info["is_multimodal"],
|
||||
"is_image": multimodal_info["is_image"],
|
||||
"multimodal_info": multimodal_info,
|
||||
"warnings": warnings
|
||||
}
|
||||
|
|
@ -675,7 +675,7 @@ def format_and_template_dataset(
|
|||
"final_format": "vlm_messages",
|
||||
"chat_column": "messages",
|
||||
"is_vlm": True,
|
||||
"is_multimodal": True,
|
||||
"is_image": True,
|
||||
"multimodal_info": multimodal_info,
|
||||
"success": True,
|
||||
"requires_manual_mapping": False,
|
||||
|
|
@ -797,7 +797,7 @@ def format_and_template_dataset(
|
|||
"final_format": "vlm_messages",
|
||||
"chat_column": "messages",
|
||||
"is_vlm": True,
|
||||
"is_multimodal": multimodal_info["is_multimodal"],
|
||||
"is_image": multimodal_info["is_image"],
|
||||
"multimodal_info": multimodal_info,
|
||||
"vlm_structure": vlm_structure,
|
||||
"success": True,
|
||||
|
|
@ -826,7 +826,7 @@ def format_and_template_dataset(
|
|||
# Gemma emits a leading <bos> that must be stripped for text-only chatml/sharegpt.
|
||||
is_alpaca = format_type == "alpaca" or (format_type == "auto" and dataset_info["detected_format"] == "alpaca")
|
||||
is_gemma = "gemma" in model_name.lower()
|
||||
if is_gemma and not dataset_info["is_multimodal"] and not is_alpaca:
|
||||
if is_gemma and not dataset_info["is_image"] and not is_alpaca:
|
||||
remove_bos_prefix = True
|
||||
template_result = apply_chat_template_to_dataset(
|
||||
dataset_info=dataset_info,
|
||||
|
|
|
|||
|
|
@ -334,7 +334,7 @@ def detect_multimodal_dataset(dataset):
|
|||
|
||||
Returns:
|
||||
dict: {
|
||||
"is_multimodal": bool,
|
||||
"is_image": bool,
|
||||
"multimodal_columns": list of column names containing image data,
|
||||
"modality_types": list of detected types (e.g., ["image", "audio"]),
|
||||
"is_audio": bool,
|
||||
|
|
@ -427,7 +427,7 @@ def detect_multimodal_dataset(dataset):
|
|||
break
|
||||
|
||||
return {
|
||||
"is_multimodal": len(multimodal_columns) > 0 or is_audio,
|
||||
"is_image": len(multimodal_columns) > 0,
|
||||
"multimodal_columns": multimodal_columns,
|
||||
"modality_types": list(modality_types),
|
||||
"is_audio": is_audio,
|
||||
|
|
|
|||
|
|
@ -423,14 +423,13 @@ def safe_num_proc(desired: Optional[int] = None) -> int:
|
|||
"""
|
||||
Return a safe ``num_proc`` for ``dataset.map()`` calls.
|
||||
|
||||
Fork-based multiprocessing deadlocks when CUDA has already been
|
||||
initialized (e.g. after inference). This helper detects that case
|
||||
and forces ``num_proc=1``.
|
||||
|
||||
On multi-GPU machines the NVIDIA driver spawns extra background threads,
|
||||
making ``os.fork()`` prone to deadlocks when many workers are created.
|
||||
This helper caps ``num_proc`` to 4 on such machines.
|
||||
|
||||
On single-GPU (or CPU-only) machines the original value is returned
|
||||
unchanged.
|
||||
|
||||
Args:
|
||||
desired: The num_proc you *want*. If None, auto-computes from
|
||||
``os.cpu_count()``.
|
||||
|
|
@ -443,14 +442,6 @@ def safe_num_proc(desired: Optional[int] = None) -> int:
|
|||
if desired is None or not isinstance(desired, int):
|
||||
desired = max(1, os.cpu_count() // 3)
|
||||
|
||||
# After inference, CUDA is initialized — forking will deadlock.
|
||||
try:
|
||||
import torch
|
||||
if torch.cuda.is_initialized():
|
||||
return 1
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
if get_physical_gpu_count() > 1:
|
||||
capped = min(4, desired)
|
||||
print(
|
||||
|
|
|
|||
|
|
@ -62,10 +62,10 @@ export function DatasetPreviewDialog({
|
|||
);
|
||||
const { isStarting, startError, startTrainingRun } = useTrainingActions();
|
||||
|
||||
// If the backend reports multimodal data, treat as VLM even if the prop
|
||||
// hasn't caught up yet (isDatasetMultimodal may still be null in the store).
|
||||
// If the backend reports image data, treat as VLM even if the prop
|
||||
// hasn't caught up yet (isDatasetImage may still be null in the store).
|
||||
const effectiveIsAudio = !!data?.is_audio;
|
||||
const effectiveIsVlm = !effectiveIsAudio && (isVlm || !!data?.is_multimodal);
|
||||
const effectiveIsVlm = isVlm || !!data?.is_image;
|
||||
|
||||
const hasHeuristicMapping = !data?.requires_manual_mapping && !!data?.suggested_mapping;
|
||||
const mappingEnabled = !!data?.requires_manual_mapping || hasHeuristicMapping;
|
||||
|
|
|
|||
|
|
@ -114,7 +114,7 @@ function SliderRow({
|
|||
export function ParamsSection(): ReactElement {
|
||||
const store = useTrainingConfigStore();
|
||||
const isLora = store.trainingMethod !== "full";
|
||||
const showVisionLora = store.isVisionModel && store.isDatasetMultimodal === true;
|
||||
const showVisionLora = store.isVisionModel && store.isDatasetImage === true;
|
||||
const [loraOpen, setLoraOpen] = useState(false);
|
||||
const [hyperOpen, setHyperOpen] = useState(false);
|
||||
const maxStepsSliderMax = Math.max(500, store.maxSteps, 30);
|
||||
|
|
|
|||
|
|
@ -42,7 +42,7 @@ export function TrainingSection() {
|
|||
const store = useTrainingConfigStore();
|
||||
const { isStarting, startError, startTrainingRun } = useTrainingActions();
|
||||
const isIncompatible =
|
||||
!store.isVisionModel && !store.isDatasetAudio && store.isDatasetMultimodal === true;
|
||||
!store.isVisionModel && store.isDatasetImage === true;
|
||||
const fileInputRef = useRef<HTMLInputElement>(null);
|
||||
|
||||
const handleFileUpload = (e: React.ChangeEvent<HTMLInputElement>) => {
|
||||
|
|
|
|||
|
|
@ -93,7 +93,7 @@ export function StudioPage(): ReactElement {
|
|||
datasetSplit={config.datasetSplit}
|
||||
mode={dialogMode}
|
||||
initialData={dialogInitial}
|
||||
isVlm={config.isVisionModel && config.isDatasetMultimodal === true}
|
||||
isVlm={config.isVisionModel && config.isDatasetImage === true}
|
||||
/>
|
||||
|
||||
{canGoBack && (
|
||||
|
|
|
|||
|
|
@ -57,7 +57,7 @@ export function buildTrainingStartPayload(
|
|||
finetune_language_layers: config.finetuneLanguageLayers,
|
||||
finetune_attention_modules: config.finetuneAttentionModules,
|
||||
finetune_mlp_modules: config.finetuneMLPModules,
|
||||
is_dataset_multimodal: !!config.isDatasetMultimodal,
|
||||
is_dataset_image: !!config.isDatasetImage,
|
||||
is_dataset_audio: config.isDatasetAudio,
|
||||
enable_wandb: config.enableWandb,
|
||||
wandb_token: config.enableWandb ? config.wandbToken.trim() || null : null,
|
||||
|
|
|
|||
|
|
@ -36,7 +36,7 @@ export function useTrainingActions() {
|
|||
|
||||
try {
|
||||
const datasetName = getDatasetName(config);
|
||||
let isVlm = config.isVisionModel && config.isDatasetMultimodal === true;
|
||||
let isVlm = config.isVisionModel && config.isDatasetImage === true;
|
||||
|
||||
if (datasetName) {
|
||||
const check = await checkDatasetFormat({
|
||||
|
|
@ -47,17 +47,17 @@ export function useTrainingActions() {
|
|||
isVlm,
|
||||
});
|
||||
|
||||
// Backend auto-detects multimodal/audio from dataset content.
|
||||
// Backend auto-detects image/audio from dataset content.
|
||||
// Sync these flags into the store so buildTrainingStartPayload picks them up.
|
||||
const isAudio = !!check.is_audio;
|
||||
const isMultimodal = !!check.is_multimodal;
|
||||
const isImage = !!check.is_image;
|
||||
|
||||
if (isMultimodal && config.isVisionModel) {
|
||||
if (isImage && config.isVisionModel) {
|
||||
isVlm = true;
|
||||
}
|
||||
if (isMultimodal !== config.isDatasetMultimodal || isAudio !== config.isDatasetAudio) {
|
||||
if (isImage !== config.isDatasetImage || isAudio !== config.isDatasetAudio) {
|
||||
useTrainingConfigStore.setState({
|
||||
isDatasetMultimodal: isMultimodal,
|
||||
isDatasetImage: isImage,
|
||||
isDatasetAudio: isAudio,
|
||||
});
|
||||
}
|
||||
|
|
|
|||
|
|
@ -35,7 +35,7 @@ const initialState: TrainingConfigState = {
|
|||
modelDefaultsError: null,
|
||||
modelDefaultsAppliedFor: null,
|
||||
isCheckingDataset: false,
|
||||
isDatasetMultimodal: null,
|
||||
isDatasetImage: null,
|
||||
isDatasetAudio: false,
|
||||
...DEFAULT_HYPERPARAMS,
|
||||
};
|
||||
|
|
@ -57,7 +57,7 @@ const NON_PERSISTED_STATE_KEYS: ReadonlySet<keyof TrainingConfigState> = new Set
|
|||
"modelDefaultsError",
|
||||
"modelDefaultsAppliedFor",
|
||||
"isCheckingDataset",
|
||||
"isDatasetMultimodal",
|
||||
"isDatasetImage",
|
||||
"isDatasetAudio",
|
||||
"trainOnCompletions",
|
||||
]);
|
||||
|
|
@ -116,9 +116,9 @@ export const useTrainingConfigStore = create<TrainingConfigStore>()(
|
|||
_trainOnCompletionsManuallySet = false;
|
||||
const patch = mapBackendModelConfigToTrainingPatch(modelDetails.config);
|
||||
|
||||
// If vision model + multimodal dataset already known, override
|
||||
// If vision model + image dataset already known, override
|
||||
// trainOnCompletions to false regardless of backend default.
|
||||
if (modelDetails.is_vision && get().isDatasetMultimodal === true) {
|
||||
if (modelDetails.is_vision && get().isDatasetImage === true) {
|
||||
patch.trainOnCompletions = false;
|
||||
}
|
||||
|
||||
|
|
@ -174,16 +174,16 @@ export const useTrainingConfigStore = create<TrainingConfigStore>()(
|
|||
})
|
||||
.then((res) => {
|
||||
if (controller.signal.aborted) return;
|
||||
const isMultimodal = !!res.is_multimodal;
|
||||
const isImage = !!res.is_image;
|
||||
const isAudio = !!res.is_audio;
|
||||
const updates: Record<string, unknown> = {
|
||||
isDatasetMultimodal: isMultimodal,
|
||||
isDatasetImage: isImage,
|
||||
isDatasetAudio: isAudio,
|
||||
isCheckingDataset: false,
|
||||
};
|
||||
if (!_trainOnCompletionsManuallySet) {
|
||||
const { isVisionModel } = get();
|
||||
if (isVisionModel && isMultimodal) {
|
||||
if (isVisionModel && isImage) {
|
||||
updates.trainOnCompletions = false;
|
||||
}
|
||||
}
|
||||
|
|
@ -191,7 +191,7 @@ export const useTrainingConfigStore = create<TrainingConfigStore>()(
|
|||
})
|
||||
.catch(() => {
|
||||
if (controller.signal.aborted) return;
|
||||
set({ isDatasetMultimodal: null, isCheckingDataset: false });
|
||||
set({ isDatasetImage: null, isCheckingDataset: false });
|
||||
});
|
||||
};
|
||||
|
||||
|
|
@ -261,7 +261,7 @@ export const useTrainingConfigStore = create<TrainingConfigStore>()(
|
|||
datasetSplit: null,
|
||||
datasetEvalSplit: null,
|
||||
datasetManualMapping: emptyManualMapping(),
|
||||
isDatasetMultimodal: null,
|
||||
isDatasetImage: null,
|
||||
isCheckingDataset: false,
|
||||
});
|
||||
},
|
||||
|
|
@ -274,7 +274,7 @@ export const useTrainingConfigStore = create<TrainingConfigStore>()(
|
|||
datasetSplit: null,
|
||||
datasetEvalSplit: null,
|
||||
datasetManualMapping: emptyManualMapping(),
|
||||
isDatasetMultimodal: null,
|
||||
isDatasetImage: null,
|
||||
isCheckingDataset: false,
|
||||
});
|
||||
},
|
||||
|
|
@ -282,7 +282,7 @@ export const useTrainingConfigStore = create<TrainingConfigStore>()(
|
|||
set({
|
||||
datasetSplit,
|
||||
datasetManualMapping: emptyManualMapping(),
|
||||
isDatasetMultimodal: null,
|
||||
isDatasetImage: null,
|
||||
isCheckingDataset: false,
|
||||
});
|
||||
|
||||
|
|
@ -298,7 +298,7 @@ export const useTrainingConfigStore = create<TrainingConfigStore>()(
|
|||
ensureDatasetChecked: () => {
|
||||
const state = get();
|
||||
if (state.isCheckingDataset) return;
|
||||
if (state.isDatasetMultimodal !== null) return;
|
||||
if (state.isDatasetImage !== null) return;
|
||||
|
||||
const datasetName =
|
||||
state.datasetSource === "huggingface"
|
||||
|
|
|
|||
|
|
@ -38,7 +38,7 @@ export interface TrainingStartRequest {
|
|||
finetune_language_layers: boolean;
|
||||
finetune_attention_modules: boolean;
|
||||
finetune_mlp_modules: boolean;
|
||||
is_dataset_multimodal: boolean;
|
||||
is_dataset_image: boolean;
|
||||
is_dataset_audio: boolean;
|
||||
enable_wandb: boolean;
|
||||
wandb_token: string | null;
|
||||
|
|
|
|||
|
|
@ -59,7 +59,7 @@ export interface TrainingConfigState {
|
|||
modelDefaultsError: string | null;
|
||||
modelDefaultsAppliedFor: string | null;
|
||||
isCheckingDataset: boolean;
|
||||
isDatasetMultimodal: boolean | null;
|
||||
isDatasetImage: boolean | null;
|
||||
isDatasetAudio: boolean;
|
||||
finetuneVisionLayers: boolean;
|
||||
finetuneLanguageLayers: boolean;
|
||||
|
|
|
|||
|
|
@ -9,7 +9,7 @@ export type CheckFormatResponse = {
|
|||
detected_speaker_column?: string | null;
|
||||
preview_samples?: Record<string, unknown>[] | null;
|
||||
total_rows?: number | null;
|
||||
is_multimodal?: boolean;
|
||||
is_image?: boolean;
|
||||
is_audio?: boolean;
|
||||
multimodal_columns?: string[] | null;
|
||||
};
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue