Fix/studio full finetuning (#4391)

* Wire Studio full finetuning into training loaders

* Preserve load_model positional compatibility
This commit is contained in:
DoubleMathew 2026-03-17 22:47:26 -05:00 committed by GitHub
commit fd72376a7e
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 21 additions and 9 deletions

View file

@ -466,6 +466,7 @@ class UnslothTrainer:
is_dataset_image: bool = False,
is_dataset_audio: bool = False,
trust_remote_code: bool = False,
full_finetuning: bool = False,
) -> bool:
"""Load model for training (supports both text and vision models)"""
self.load_in_4bit = load_in_4bit # Store for training_meta.json
@ -612,6 +613,7 @@ class UnslothTrainer:
dtype = None,
auto_model = CsmForConditionalGeneration,
load_in_4bit = False,
full_finetuning = full_finetuning,
token = hf_token,
trust_remote_code = trust_remote_code,
)
@ -626,6 +628,7 @@ class UnslothTrainer:
model_name = model_name,
dtype = None,
load_in_4bit = False,
full_finetuning = full_finetuning,
auto_model = WhisperForConditionalGeneration,
whisper_language = "English",
whisper_task = "transcribe",
@ -646,6 +649,7 @@ class UnslothTrainer:
max_seq_length = max_seq_length,
dtype = None,
load_in_4bit = load_in_4bit,
full_finetuning = full_finetuning,
token = hf_token,
trust_remote_code = trust_remote_code,
)
@ -684,6 +688,7 @@ class UnslothTrainer:
max_seq_length = max_seq_length,
dtype = torch.float32, # Spark-TTS requires float32
load_in_4bit = False,
full_finetuning = full_finetuning,
token = hf_token,
trust_remote_code = trust_remote_code,
)
@ -697,6 +702,7 @@ class UnslothTrainer:
model_name,
max_seq_length = max_seq_length,
load_in_4bit = False,
full_finetuning = full_finetuning,
token = hf_token,
trust_remote_code = trust_remote_code,
)
@ -712,6 +718,7 @@ class UnslothTrainer:
max_seq_length = max_seq_length,
dtype = None,
load_in_4bit = load_in_4bit,
full_finetuning = full_finetuning,
token = hf_token,
trust_remote_code = trust_remote_code,
)
@ -724,6 +731,7 @@ class UnslothTrainer:
max_seq_length = max_seq_length,
dtype = None, # Auto-detect
load_in_4bit = load_in_4bit,
full_finetuning = full_finetuning,
token = hf_token,
trust_remote_code = trust_remote_code,
)
@ -755,6 +763,7 @@ class UnslothTrainer:
max_seq_length = max_seq_length,
dtype = None, # Auto-detect
load_in_4bit = load_in_4bit,
full_finetuning = full_finetuning,
token = hf_token,
trust_remote_code = trust_remote_code,
)
@ -779,13 +788,14 @@ class UnslothTrainer:
self._source_code_retried = True
logger.info(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,
trust_remote_code,
model_name = model_name,
max_seq_length = max_seq_length,
load_in_4bit = load_in_4bit,
hf_token = hf_token,
is_dataset_image = is_dataset_image,
is_dataset_audio = is_dataset_audio,
trust_remote_code = trust_remote_code,
full_finetuning = full_finetuning,
)
error_msg = str(e)
error_lower = error_msg.lower()

View file

@ -444,12 +444,16 @@ def run_training_process(
_tqdm_thread = _th.Thread(target = _monitor_tqdm, daemon = True)
_tqdm_thread.start()
training_type = config.get("training_type", "LoRA/QLoRA")
use_lora = training_type == "LoRA/QLoRA"
# ── 4c. Load training model (uses VRAM — dataset already formatted) ──
_send_status(event_queue, "Loading model...")
success = trainer.load_model(
model_name = model_name,
max_seq_length = config["max_seq_length"],
load_in_4bit = config["load_in_4bit"],
full_finetuning = not use_lora,
hf_token = hf_token,
is_dataset_image = config.get("is_dataset_image", False),
is_dataset_audio = config.get("is_dataset_audio", False),
@ -473,8 +477,6 @@ def run_training_process(
return
# ── 4d. Prepare model (LoRA or full finetuning) ──
training_type = config.get("training_type", "LoRA/QLoRA")
use_lora = training_type == "LoRA/QLoRA"
if use_lora:
_send_status(event_queue, "Configuring LoRA adapters...")
success = trainer.prepare_model_for_training(