FannyFa-Model-V1 / trainer.py
FannyFa's picture
Upload 16 files
3eecd6b verified
Raw History Blame Contribute Delete
64.5 kB
import os
import re
import time
import math
import gc
import json
import shutil
from pathlib import Path
from typing import List, Dict, Any, Optional, Tuple
from inspect import signature
import torch
import torch.distributed as dist
from transformers import (
AutoTokenizer,
AutoModelForCausalLM,
Trainer,
TrainingArguments,
DataCollatorForLanguageModeling,
TrainerCallback,
set_seed,
)
try:
from peft import (
LoraConfig,
TaskType,
get_peft_model,
prepare_model_for_kbit_training,
PeftModel,
)
HAS_PEFT = True
except ImportError:
HAS_PEFT = False
LoraConfig = None
TaskType = None
get_peft_model = None
prepare_model_for_kbit_training = None
PeftModel = None
# ============================================================================
# WORKAROUND: peft + torchao version mismatch
# ============================================================================
def _patch_peft_torchao_check() -> None:
"""
Force-disable PEFT's torchao availability check.
PEFT versi baru (>=0.14) raise ImportError saat is_torchao_available()
mendeteksi torchao < 0.16.0. LoRA biasa tidak butuh torchao, jadi kita
matikan check-nya pakai multi-strategy:
1. Override importlib.metadata.version("torchao") → "0.16.0"
2. Override torchao.__version__ kalau sudah di-import
3. Override is_torchao_available di semua modul peft.*
"""
import sys
import importlib.metadata as _md
# ------------------------------------------------------------------
# Strategy 1: intercept importlib.metadata.version (paling robust)
# PEFT versi baru cek versi lewat importlib.metadata.version("torchao")
# ------------------------------------------------------------------
if not getattr(_md, "_torchao_patched", False):
_orig_version = _md.version
def _patched_version(name: str):
try:
v = _orig_version(name)
except Exception:
return v
if name.lower() == "torchao":
# Fake version supaya lolos check ">= 0.16.0"
return "0.16.0"
return v
_md.version = _patched_version
_md._torchao_patched = True
_md._torchao_orig_version = _orig_version
# ------------------------------------------------------------------
# Strategy 2: override torchao.__version__ kalau sudah di-import
# ------------------------------------------------------------------
if "torchao" in sys.modules:
try:
sys.modules["torchao"].__version__ = "0.16.0"
except Exception:
pass
# ------------------------------------------------------------------
# Strategy 3: override is_torchao_available di semua modul peft.*
# ------------------------------------------------------------------
def _always_false():
return False
try:
import peft.import_utils as _iu
fn = getattr(_iu, "is_torchao_available", None)
if fn is not None and hasattr(fn, "cache_clear"):
try:
fn.cache_clear()
except Exception:
pass
_iu.is_torchao_available = _always_false
for attr in ("_torchao_available", "_is_torchao_available"):
if hasattr(_iu, attr):
try:
setattr(_iu, attr, False)
except Exception:
pass
except Exception:
pass
# Override di semua submodule peft.* yang mungkin sudah import
for mod_name, mod in list(sys.modules.items()):
if not mod_name.startswith("peft"):
continue
if getattr(mod, "is_torchao_available", None) is not None:
try:
mod.is_torchao_available = _always_false
except Exception:
pass
for attr in ("_torchao_available", "_is_torchao_available"):
if hasattr(mod, attr):
try:
setattr(mod, attr, False)
except Exception:
pass
if HAS_PEFT:
_patch_peft_torchao_check()
from datasets import Dataset
try:
import bitsandbytes as bnb
HAS_BNB = True
except ImportError:
HAS_BNB = False
from rich.markup import escape as rich_escape
from config import (
MODEL_DIR,
CHECKPOINT_DIR,
TEMP_DIR,
DEFAULT_MODEL,
DEFAULT_BATCH_SIZE,
DEFAULT_LEARNING_RATE,
DEFAULT_EPOCHS,
DEFAULT_SEQ_LEN,
DEFAULT_WARMUP_RATIO,
DEFAULT_MEMORY_LIMIT_GB,
)
from utils import (
console,
Theme,
timer,
debug_logger,
error_logger,
train_logger,
MemoryWatchdog,
set_memory_hard_limit,
graceful_exit,
get_device,
get_gpu_info,
)
from data_loader import EnhancedDatasetLoader, DatasetStats
from models import ModelManager, ModelMetadata
from history import TrainingHistory
from chat import EnhancedChatModule
from chat_template import tokenize_batch_for_training, detect_format
class EnhancedCallback(TrainerCallback):
def __init__(
self,
history: Optional[TrainingHistory] = None,
cleanup_interval: int = 100,
memory_watchdog: Optional[MemoryWatchdog] = None,
log_interval: int = 10,
):
self.history = history or TrainingHistory()
self.cleanup_interval = cleanup_interval
self.watchdog = memory_watchdog
self.start_time = time.time()
self.total_tokens = 0
self.log_interval = log_interval
self.step_times: List[float] = []
self.last_log_time = time.time()
self._overfit_check_counter = 0
self._last_step = 0
def on_step_end(self, args, state, control, **kwargs) -> None:
if state.global_step % self.cleanup_interval == 0:
if torch.cuda.is_available():
torch.cuda.empty_cache()
gc.collect()
current_time = time.time()
if state.global_step > 0:
self.step_times.append(current_time - self.last_log_time)
self.last_log_time = current_time
self._last_step = state.global_step
def on_log(self, args, state, control, logs=None, **kwargs) -> None:
if logs is None:
return
step = state.global_step
train_loss = logs.get("loss")
eval_loss = logs.get("eval_loss")
lr = logs.get("learning_rate")
if eval_loss is not None and eval_loss < 1000 and eval_loss > 0:
try:
perplexity = math.exp(eval_loss)
except OverflowError:
perplexity = float("inf")
else:
perplexity = None
accuracy = logs.get("accuracy")
grad_norm = logs.get("grad_norm")
elapsed = time.time() - self.start_time
throughput = state.global_step / max(elapsed, 1) if state.global_step > 0 else 0
memory_usage = None
if self.watchdog:
stats = self.watchdog.get_stats()
memory_usage = stats.get("rss_gb")
self.history.append(
step=step,
train_loss=train_loss,
eval_loss=eval_loss,
lr=lr,
perplexity=perplexity,
accuracy=accuracy,
grad_norm=grad_norm,
throughput=throughput,
memory_usage=memory_usage,
epoch=state.epoch,
)
def on_train_end(self, args, state, control, **kwargs) -> None:
if torch.cuda.is_available():
torch.cuda.empty_cache()
gc.collect()
class SafeEarlyStoppingCallback(TrainerCallback):
def __init__(
self, early_stopping_patience: int = 3, early_stopping_threshold: float = 0.001
):
super().__init__()
self.patience = early_stopping_patience
self.threshold = early_stopping_threshold
self.best_score = None
self.best_step = None
self.wait = 0
self.metric = None
self._initialized = False
def on_evaluate(self, args, state, control, metrics, **kwargs) -> None:
if not self._initialized:
self.metric = getattr(args, "metric_for_best_model", "eval_loss")
self._initialized = True
current_score = metrics.get(self.metric)
if current_score is None:
current_score = metrics.get("eval_loss")
if current_score is None:
return
if self.best_score is None or current_score < self.best_score - self.threshold:
self.best_score = current_score
self.best_step = state.global_step
self.wait = 0
else:
self.wait += 1
if self.wait >= self.patience:
control.should_training_stop = True
console.print(
Theme.success(
f" Early stopping at step {state.global_step} "
f"(best {self.metric}: {self.best_score:.4f} at step {self.best_step})"
)
)
class GradientAccumulationScheduler(TrainerCallback):
def __init__(
self,
initial_steps: int = 1,
max_steps: int = 16,
min_steps: int = 1,
memory_threshold: float = 0.7,
adjustment_interval: int = 100,
):
self.initial_steps = initial_steps
self.max_steps = max_steps
self.min_steps = min_steps
self.memory_threshold = memory_threshold
self.adjustment_interval = adjustment_interval
self.current_steps = initial_steps
self._last_adjustment = 0
def on_step_end(self, args, state, control, **kwargs) -> None:
if state.global_step % self.adjustment_interval != 0:
return
if not torch.cuda.is_available():
return
try:
allocated = torch.cuda.memory_allocated() / (1024**3)
total = torch.cuda.get_device_properties(0).total_memory / (1024**3)
memory_ratio = allocated / total
if (
memory_ratio > self.memory_threshold
and self.current_steps < self.max_steps
):
new_steps = min(self.current_steps * 2, self.max_steps)
if new_steps != self.current_steps:
self.current_steps = new_steps
args.gradient_accumulation_steps = self.current_steps
debug_logger.debug(
f"Grad accumulation increased to {self.current_steps}"
)
elif (
memory_ratio < self.memory_threshold * 0.4
and self.current_steps > self.min_steps
):
new_steps = max(self.current_steps // 2, self.min_steps)
if new_steps != self.current_steps:
self.current_steps = new_steps
args.gradient_accumulation_steps = self.current_steps
debug_logger.debug(
f"Grad accumulation decreased to {self.current_steps}"
)
except Exception as e:
debug_logger.debug(f"Grad accumulation scheduler error: {e}")
def make_training_args_safe(
checkpoint_dir: str, train_cfg: Dict, device: str, multi_gpu: bool = False
) -> TrainingArguments:
sig = signature(TrainingArguments.__init__)
supported = set(sig.parameters.keys()) - {"self", "args", "kwargs"}
if "eval_strategy" in supported:
EVAL_STRAT_KEY = "eval_strategy"
elif "evaluation_strategy" in supported:
EVAL_STRAT_KEY = "evaluation_strategy"
else:
EVAL_STRAT_KEY = None
kwargs = {
"output_dir": checkpoint_dir,
"num_train_epochs": train_cfg.get("num_epochs", DEFAULT_EPOCHS),
"per_device_train_batch_size": train_cfg.get("batch_size", DEFAULT_BATCH_SIZE),
"per_device_eval_batch_size": train_cfg.get("batch_size", DEFAULT_BATCH_SIZE),
"learning_rate": train_cfg.get("learning_rate", DEFAULT_LEARNING_RATE),
"seed": train_cfg.get("seed", 42),
"report_to": "none",
}
if multi_gpu:
if "ddp_find_unused_parameters" in supported:
kwargs["ddp_find_unused_parameters"] = False
if "ddp_bucket_cap_mb" in supported:
kwargs["ddp_bucket_cap_mb"] = 25
optional_args = {
"save_steps": "save_steps",
"logging_steps": "logging_steps",
"warmup_steps": "warmup_steps",
"weight_decay": "weight_decay",
"gradient_accumulation_steps": "gradient_accumulation_steps",
"fp16": "fp16",
"bf16": "bf16",
"eval_steps": "eval_steps",
"save_total_limit": "save_total_limit",
"gradient_clip_value": "gradient_clip_value",
"load_best_model_at_end": "load_best_model_at_end",
"metric_for_best_model": "metric_for_best_model",
"greater_is_better": "greater_is_better",
"dataloader_num_workers": "dataloader_num_workers",
"optim": "optim",
"lr_scheduler_type": "lr_scheduler_type",
"max_grad_norm": "max_grad_norm",
"label_smoothing_factor": "label_smoothing_factor",
"group_by_length": "group_by_length",
"disable_tqdm": "disable_tqdm",
}
for cfg_key, arg_name in optional_args.items():
if cfg_key in train_cfg and arg_name in supported:
kwargs[arg_name] = train_cfg[cfg_key]
if "warmup_ratio" in train_cfg and "warmup_ratio" in supported:
kwargs["warmup_ratio"] = train_cfg["warmup_ratio"]
if kwargs.get("fp16") and device != "cuda":
kwargs["fp16"] = False
if kwargs.get("bf16") and device != "cuda":
kwargs["bf16"] = False
early_stopping_patience = train_cfg.get("early_stopping_patience", 0)
load_best = train_cfg.get("load_best_model_at_end", False)
if early_stopping_patience > 0:
kwargs["load_best_model_at_end"] = True
if "metric_for_best_model" not in kwargs:
kwargs["metric_for_best_model"] = "eval_loss"
if "greater_is_better" not in kwargs:
kwargs["greater_is_better"] = False
if EVAL_STRAT_KEY is not None:
if "save_steps" in kwargs:
kwargs[EVAL_STRAT_KEY] = "steps"
if "eval_steps" not in kwargs:
kwargs["eval_steps"] = kwargs.get("save_steps", 500)
else:
kwargs[EVAL_STRAT_KEY] = "epoch"
if "save_strategy" in supported:
kwargs["save_strategy"] = kwargs[EVAL_STRAT_KEY]
elif load_best and EVAL_STRAT_KEY is not None:
kwargs[EVAL_STRAT_KEY] = "steps" if "save_steps" in kwargs else "epoch"
if "save_strategy" in supported:
kwargs["save_strategy"] = kwargs[EVAL_STRAT_KEY]
elif EVAL_STRAT_KEY is not None:
kwargs[EVAL_STRAT_KEY] = "no"
try:
return TrainingArguments(**kwargs)
except TypeError as e:
debug_logger.debug(f"TrainingArguments fallback: {e}")
fallback = {
"output_dir": checkpoint_dir,
"num_train_epochs": train_cfg.get("num_epochs", DEFAULT_EPOCHS),
"per_device_train_batch_size": train_cfg.get(
"batch_size", DEFAULT_BATCH_SIZE
),
"learning_rate": train_cfg.get("learning_rate", DEFAULT_LEARNING_RATE),
"seed": train_cfg.get("seed", 42),
}
if EVAL_STRAT_KEY is not None:
fallback[EVAL_STRAT_KEY] = "no"
return TrainingArguments(**fallback)
def compute_bleu_rouge(
predictions: List[str], references: List[str]
) -> Dict[str, float]:
results = {}
if not predictions or not references:
return results
min_len = min(len(predictions), len(references))
predictions = predictions[:min_len]
references = references[:min_len]
total = 0.0
for pred, ref in zip(predictions, references):
pred_tokens = set(re.findall(r"\w+", pred.lower()))
ref_tokens = set(re.findall(r"\w+", ref.lower()))
if ref_tokens:
total += len(pred_tokens & ref_tokens) / len(ref_tokens)
results["unigram_overlap"] = total / len(predictions) if predictions else 0.0
try:
import evaluate
HAS_EVALUATE = True
except ImportError:
HAS_EVALUATE = False
if HAS_EVALUATE:
try:
bleu = evaluate.load("bleu")
bleu_result = bleu.compute(
predictions=predictions, references=[[r] for r in references]
)
results["bleu"] = bleu_result.get("bleu", 0.0)
except Exception as e:
debug_logger.debug(f"BLEU evaluate failed: {e}")
try:
rouge = evaluate.load("rouge")
rouge_result = rouge.compute(predictions=predictions, references=references)
for key, val in rouge_result.items():
if "rougeL" in key.lower():
results["rougeL"] = val
break
if "rougeL" not in results and rouge_result:
results["rouge"] = next(iter(rouge_result.values()))
except Exception as e:
debug_logger.debug(f"ROUGE evaluate failed: {e}")
try:
import sacrebleu
HAS_SACREBLEU = True
except ImportError:
HAS_SACREBLEU = False
if "bleu" not in results and HAS_SACREBLEU:
try:
bleu_score = sacrebleu.corpus_bleu(predictions, [references])
results["bleu_sacrebleu"] = bleu_score.score / 100
except Exception as e:
debug_logger.debug(f"sacrebleu failed: {e}")
return results
# ---------------------------------------------------------------------------
# Top-level, picklable helper functions untuk dataset.map / .filter
# ---------------------------------------------------------------------------
def _tokenize_batch(examples, tokenizer, max_length):
"""Tokenize dengan prompt masking (auto-detect format Qwen vs GPT-2)."""
return tokenize_batch_for_training(examples, tokenizer, max_length)
def _filter_valid_lengths(examples, min_len, max_len):
return [
len(ids) >= min_len and len(ids) <= max_len
for ids in examples["input_ids"]
]
class EnhancedTrainingModule:
def __init__(self):
self.model_manager = ModelManager()
self.device = get_device()
self.is_multi_gpu = self.model_manager.is_multi_gpu()
self.tokenizer = None
self.model = None
self.is_peft_model = False
self.watchdog = None
self.history = TrainingHistory()
self._current_step = 0
self._start_time = None
console.print(Theme.dim(f"Device: {self.device.upper()}"))
if self.device == "cuda":
gpu_info = get_gpu_info()
if gpu_info:
console.print(
Theme.dim(
f"GPU: {gpu_info.get('name', 'Unknown')} ({gpu_info.get('memory_total_gb', 0):.1f} GB)"
)
)
def run(self) -> None:
from rich.panel import Panel
from rich.prompt import Prompt, Confirm
console.print(
Panel(
Theme.header(" ENHANCED TRAINING MODULE v3")
+ "\n"
+ Theme.dim(
"Auto-LoRA • 8-bit • Multi-GPU • Qwen/GPT-2 auto-detect"
),
title="TRAINING",
style="bold yellow",
)
)
self._setup_memory_protection()
console.print(Theme.info("Pilihan:"))
console.print(" [green]1. Training baru[/green]")
console.print(" [green]2. Training dengan dataset besar (streaming)[/green]")
console.print(" [green]3. Resume training dari checkpoint[/green]")
console.print(" [green]4. Advanced Training (PEFT/LoRA)[/green]")
console.print(" [green]5. Kembali[/green]")
choice = Prompt.ask("[yellow]Pilih", choices=["1", "2", "3", "4", "5"])
if choice == "1":
self._train_new(streaming=False)
elif choice == "2":
self._train_new(streaming=True)
elif choice == "3":
self._resume_training()
elif choice == "4":
self._train_advanced()
else:
self._cleanup()
return
def _setup_memory_protection(self) -> None:
mem_limit = self.model_manager.config.get(
"memory_limit_gb", DEFAULT_MEMORY_LIMIT_GB
)
if mem_limit > 0:
set_memory_hard_limit(mem_limit)
self.watchdog = MemoryWatchdog(
limit_gb=mem_limit, step_getter=lambda: self._current_step
)
self.watchdog.start()
def _train_new(self, streaming: bool = False) -> None:
from rich.prompt import Prompt, Confirm
from pathlib import Path
dataset_path = self._select_dataset_file()
if dataset_path is None:
return
if not os.path.exists(dataset_path):
console.print(Theme.error("File tidak ditemukan!"))
return
console.print(Theme.info("Membaca dataset..."))
try:
loader = EnhancedDatasetLoader()
train_cfg = self.model_manager.get_training_config()
augment_cfg = train_cfg.get("data_augmentation", {})
if streaming:
console.print(Theme.warning(" Streaming mode aktif"))
all_samples = []
checkpoint_file = os.path.join(
TEMP_DIR, f"resume_{Path(dataset_path).stem}.json"
)
for chunk in loader.load_streaming_with_resume(
dataset_path, checkpoint_file
):
all_samples.extend(chunk)
if graceful_exit.should_exit:
console.print(Theme.warning(" Loading interrupted"))
return
samples = all_samples
stats = DatasetStats()
stats.total_samples = len(samples)
stats.valid_samples = len(samples)
stats.format_detected = Path(dataset_path).suffix[1:].upper()
else:
samples, stats = loader.load(
dataset_path,
augment=augment_cfg.get("enabled", False),
augment_cfg=augment_cfg,
)
except Exception as e:
console.print(Theme.error(f"Error: {rich_escape(str(e))}"))
error_logger.error(f"Dataset loading error: {e}")
return
self._show_dataset_stats(stats)
if stats.valid_samples < 10:
console.print(Theme.error("Dataset terlalu kecil! Minimal 10 sample."))
return
train_cfg = self.model_manager.get_training_config()
console.print("\n" + Theme.info("Konfigurasi training saat ini:"))
for k, v in train_cfg.items():
if k not in ["peft_config", "data_augmentation"]:
console.print(f" {k}: {v}")
if not Confirm.ask("[yellow]Gunakan konfigurasi ini?", default=True):
train_cfg = self._customize_training_config(train_cfg)
self.model_manager.update_training_config(**train_cfg)
model_name = Prompt.ask("[cyan]Nama model", default="my-qwen-finetune")
model_path = os.path.join(MODEL_DIR, model_name)
if os.path.exists(model_path):
if not Confirm.ask(
Theme.error("Model sudah ada. Overwrite?"), default=False
):
console.print(Theme.warning("Training dibatalkan."))
return
shutil.rmtree(model_path)
self._execute_training(
samples=samples,
model_name=model_name,
dataset_path=dataset_path,
train_cfg=train_cfg,
stats=stats,
streaming=streaming,
)
def _train_advanced(self) -> None:
from rich.panel import Panel
from rich.prompt import Prompt, Confirm
console.print(
Panel(
Theme.header(" ADVANCED TRAINING WITH PEFT/LoRA")
+ "\n"
+ Theme.dim(
"Auto-detect target modules • 8-bit quantization • Memory efficient"
),
title="ADVANCED",
style="bold magenta",
)
)
console.print(
f" PEFT/LoRA: {'[green] Available[/green]' if HAS_PEFT else '[red] Not installed[/red]'}"
)
console.print(
f" 8-bit Quantization: {'[green] Available[/green]' if HAS_BNB else '[red] Not installed[/red]'}"
)
console.print(f" Device: {self.device.upper()}")
if not HAS_PEFT:
console.print(
Theme.warning(" PEFT not installed. Install with: pip install peft")
)
if not Confirm.ask("[yellow]Lanjutkan tanpa PEFT?", default=True):
return
self._train_new(streaming=False)
def _customize_training_config(self, train_cfg: Dict) -> Dict:
from rich.prompt import Prompt, Confirm
cfg = train_cfg.copy()
console.print("\n" + Theme.warning("Pilih base model:"))
console.print(" [bold green]1. Qwen/Qwen2.5-1.5B-Instruct (RECOMMENDED)[/bold green]")
console.print(" [green]2. Qwen/Qwen2.5-0.5B-Instruct (paling kecil)[/green]")
console.print(" [green]3. cahya/gpt2-small-indonesian-522M (yang lama)[/green]")
console.print(" [green]4. gpt2 (English, base)[/green]")
console.print(" [green]5. Custom model ID[/green]")
model_choice = Prompt.ask(
"Pilih", choices=["1", "2", "3", "4", "5"], default="1"
)
model_map = {
"1": "Qwen/Qwen2.5-1.5B-Instruct",
"2": "Qwen/Qwen2.5-0.5B-Instruct",
"3": "cahya/gpt2-small-indonesian-522M",
"4": "gpt2",
}
if model_choice == "5":
cfg["base_model"] = Prompt.ask("Model ID")
else:
cfg["base_model"] = model_map.get(model_choice, "Qwen/Qwen2.5-1.5B-Instruct")
if "qwen" in cfg["base_model"].lower():
console.print(Theme.dim(" Detected Qwen model — pakai default yang cocok"))
cfg.setdefault("batch_size", 2)
cfg.setdefault("num_epochs", 3)
cfg.setdefault("max_length", 512)
cfg.setdefault("learning_rate", 5e-5)
cfg.setdefault("warmup_ratio", 0.1)
cfg.setdefault("gradient_accumulation_steps", 4)
params = [
("learning_rate", float, "Learning rate", DEFAULT_LEARNING_RATE),
("batch_size", int, "Batch size", DEFAULT_BATCH_SIZE),
("num_epochs", int, "Number of epochs", DEFAULT_EPOCHS),
("max_length", int, "Max token length", DEFAULT_SEQ_LEN),
("validation_split", float, "Validation split (0-1)", 0.1),
("early_stopping_patience", int, "Early stopping patience", 3),
("gradient_accumulation_steps", int, "Gradient accumulation", 1),
("warmup_ratio", float, "Warmup ratio", DEFAULT_WARMUP_RATIO),
("weight_decay", float, "Weight decay", 0.01),
]
for key, typ, label, default in params:
current = cfg.get(key, default)
new_val = Prompt.ask(label, default=str(current))
try:
cfg[key] = typ(new_val)
except ValueError:
console.print(Theme.warning(f"Skipped {key}"))
console.print("\n" + Theme.info("Advanced options:"))
if HAS_PEFT and Confirm.ask("Gunakan PEFT/LoRA (hemat memori)?", default=False):
cfg["use_peft"] = True
cfg["peft_config"]["r"] = int(
Prompt.ask("LoRA r", default=str(cfg["peft_config"]["r"]))
)
cfg["peft_config"]["alpha"] = int(
Prompt.ask("LoRA alpha", default=str(cfg["peft_config"]["alpha"]))
)
cfg["peft_config"]["dropout"] = float(
Prompt.ask("LoRA dropout", default=str(cfg["peft_config"]["dropout"]))
)
if HAS_BNB and Confirm.ask(
"Gunakan 8-bit quantization (hemat memori ekstra)?", default=False
):
cfg["use_8bit"] = True
if Confirm.ask("Freeze embeddings?", default=False):
cfg["freeze_embeddings"] = True
if torch.cuda.device_count() > 1:
cfg["multi_gpu"] = Confirm.ask(
f"Enable multi-GPU ({torch.cuda.device_count()} GPUs)?", default=False
)
self.model_manager.config["multi_gpu"] = cfg["multi_gpu"]
return cfg
def _auto_detect_lora_target_modules(self) -> List[str]:
target_modules = []
linear_modules = []
for name, module in self.model.named_modules():
if isinstance(module, torch.nn.Linear):
linear_modules.append(name)
if hasattr(self.model, "config") and hasattr(self.model.config, "model_type"):
mt = self.model.config.model_type
if mt == "gpt2":
return ["c_attn", "c_proj", "c_fc"]
if mt in ("qwen2", "qwen3", "qwen"):
return ["q_proj", "k_proj", "v_proj", "o_proj",
"gate_proj", "up_proj", "down_proj"]
if mt in ("llama", "mistral"):
return ["q_proj", "k_proj", "v_proj", "o_proj"]
for name in linear_modules:
if "attn" in name.lower() or "attention" in name.lower():
parts = name.split(".")
if parts:
module_name = parts[-1]
if module_name and module_name not in target_modules:
target_modules.append(module_name)
if not target_modules:
common_patterns = [
"q_proj", "k_proj", "v_proj", "o_proj",
"c_attn", "c_proj",
]
for pattern in common_patterns:
if any(pattern in name for name in linear_modules):
target_modules.append(pattern)
if not target_modules:
target_modules = ["c_attn", "c_proj", "c_fc"]
console.print(
Theme.dim(f" Auto-detected LoRA target modules: {target_modules}")
)
return target_modules
def _apply_peft(self, train_cfg: Dict) -> None:
if not HAS_PEFT:
console.print(Theme.warning(" PEFT not available"))
return
try:
# ----------------------------------------------------------------
# PATCH: pastikan torchao check dimatikan sebelum get_peft_model()
# ----------------------------------------------------------------
_patch_peft_torchao_check()
peft_cfg = train_cfg.get("peft_config", {})
r = peft_cfg.get("r", 8)
alpha = peft_cfg.get("alpha", 16)
dropout = peft_cfg.get("dropout", 0.05)
if train_cfg.get("use_8bit", False) and HAS_BNB:
self.model = prepare_model_for_kbit_training(self.model)
console.print(Theme.dim(" Model prepared for k-bit training"))
target_modules = peft_cfg.get("target_modules")
if target_modules == "auto" or not target_modules:
target_modules = self._auto_detect_lora_target_modules()
elif isinstance(target_modules, str):
target_modules = [target_modules]
lora_config = LoraConfig(
r=r,
lora_alpha=alpha,
target_modules=target_modules,
lora_dropout=dropout,
bias="none",
task_type=TaskType.CAUSAL_LM,
inference_mode=False,
)
self.model = get_peft_model(self.model, lora_config)
self.is_peft_model = True
if not train_cfg.get("use_8bit", False) and not self.is_multi_gpu:
self.model.to(self.device)
trainable_params = sum(
p.numel() for p in self.model.parameters() if p.requires_grad
)
total_params = sum(p.numel() for p in self.model.parameters())
console.print(
Theme.success(
f" PEFT/LoRA applied: {trainable_params:,} trainable params "
f"({trainable_params / total_params:.2%} of total)"
)
)
except Exception as e:
err_str = str(e)
console.print(Theme.error(f"PEFT setup failed: {err_str}"))
error_logger.error(f"PEFT error: {err_str}")
if "torchao" in err_str.lower():
console.print(Theme.warning(
" Hint: jalankan `pip uninstall -y torchao` ATAU "
"`pip install -U 'torchao>=0.16.0'` lalu restart runtime."
))
elif "cuda" in err_str.lower() and "memory" in err_str.lower():
console.print(Theme.warning(
" Hint: GPU OOM. Coba use_8bit=True atau batch_size lebih kecil."
))
self.is_peft_model = False
console.print(Theme.warning(
" Fallback ke full fine-tuning (tanpa LoRA). "
"Kalau OOM, kecilkan batch_size atau pakai use_8bit=True."
))
def _auto_tune_batch_size(self, train_cfg: Dict) -> None:
try:
if self.device == "cuda" and torch.cuda.is_available():
props = torch.cuda.get_device_properties(0)
vram_gb = props.total_memory / (1024**3)
model_params = sum(p.numel() for p in self.model.parameters())
model_gb = model_params * 4 / (1024**3)
suggested = max(1, int(vram_gb / (model_gb * 2.0)))
suggested = min(suggested, 32)
current = train_cfg.get("batch_size", DEFAULT_BATCH_SIZE)
if current > suggested:
console.print(
Theme.warning(
f"GPU VRAM {vram_gb:.1f}GB - menurunkan batch_size dari {current} -> {suggested}"
)
)
train_cfg["batch_size"] = max(1, suggested)
elif current < suggested and current < 8:
console.print(
Theme.dim(
f"GPU VRAM {vram_gb:.1f}GB - batch_size dapat dinaikkan ke {min(suggested, 8)}"
)
)
except Exception as e:
debug_logger.debug(f"Auto-tune batch size error: {e}")
def _setup_multi_gpu(self) -> None:
try:
if not dist.is_initialized():
world_size = torch.cuda.device_count()
console.print(Theme.info(f" Multi-GPU mode enabled: {world_size} GPUs"))
os.environ["MASTER_ADDR"] = "localhost"
os.environ["MASTER_PORT"] = str(29500 + hash(str(os.getpid())) % 10000)
dist.init_process_group("nccl", rank=0, world_size=1)
else:
console.print(Theme.dim("Multi-GPU already initialized"))
except Exception as e:
console.print(Theme.warning(f"Multi-GPU setup failed: {e}"))
self.is_multi_gpu = False
def _execute_training(
self,
samples: List[str],
model_name: str,
dataset_path: str,
train_cfg: Dict,
stats: DatasetStats,
streaming: bool = False,
) -> None:
model_path = os.path.join(MODEL_DIR, model_name)
base_model = train_cfg.get("base_model", "Qwen/Qwen2.5-1.5B-Instruct")
self._current_step = 0
self._start_time = time.time()
console.print(Theme.info(f"Loading base model: {base_model}..."))
try:
self.tokenizer = AutoTokenizer.from_pretrained(
base_model, trust_remote_code=True
)
if self.tokenizer.pad_token is None:
self.tokenizer.pad_token = self.tokenizer.eos_token
# ------------------------------------------------------------------
# Tentukan dtype training. T4 tidak support bf16.
# ------------------------------------------------------------------
use_fp16 = bool(train_cfg.get("fp16", False))
use_bf16 = bool(train_cfg.get("bf16", False))
if (
use_bf16
and torch.cuda.is_available()
and not torch.cuda.is_bf16_supported()
):
console.print(Theme.warning(
" bf16 tidak didukung GPU ini (mis. T4) — fallback ke fp16"
))
use_bf16 = False
use_fp16 = True
train_cfg["bf16"] = False
train_cfg["fp16"] = True
# Kalau user tidak set fp16/bf16, default ke fp16 di CUDA
if not use_fp16 and not use_bf16 and self.device == "cuda":
use_fp16 = True
train_cfg["fp16"] = True
console.print(Theme.dim(" Defaulting to fp16 (CUDA)"))
# ------------------------------------------------------------------
# CRITICAL FIX untuk "Attempting to unscale FP16 gradients":
#
# GradScaler BUTUH master weights fp32:
# - Full fine-tuning (tanpa PEFT) → load model fp32
# - PEFT/LoRA → base frozen, boleh fp16 (hemat memori)
# ------------------------------------------------------------------
will_use_peft = bool(train_cfg.get("use_peft", False)) and HAS_PEFT
load_kwargs = {"low_cpu_mem_usage": True, "trust_remote_code": True}
if train_cfg.get("use_8bit", False) and HAS_BNB:
console.print(Theme.warning(" Loading model in 8-bit quantization"))
load_kwargs["load_in_8bit"] = True
load_kwargs["device_map"] = "auto"
load_kwargs["llm_int8_threshold"] = 6.0
load_kwargs["torch_dtype"] = (
torch.float16 if use_fp16 else torch.float32
)
else:
load_kwargs["device_map"] = None
if will_use_peft:
# Base frozen → boleh fp16
if use_fp16:
load_kwargs["torch_dtype"] = torch.float16
elif use_bf16:
load_kwargs["torch_dtype"] = torch.bfloat16
else:
load_kwargs["torch_dtype"] = torch.float32
console.print(Theme.dim(
" PEFT mode: base model hemat memori (fp16)"
))
else:
# Full FT → WAJIB fp32 untuk GradScaler master weights
load_kwargs["torch_dtype"] = torch.float32
console.print(Theme.warning(
" Full FT: model di-load sebagai fp32 (master weights). "
"Aktifkan use_peft=True kalau VRAM tidak cukup."
))
self.model = AutoModelForCausalLM.from_pretrained(base_model, **load_kwargs)
try:
actual_dtype = next(self.model.parameters()).dtype
console.print(Theme.dim(f" Model dtype: {actual_dtype}"))
except StopIteration:
pass
if hasattr(self.model, "gradient_checkpointing_enable"):
try:
self.model.gradient_checkpointing_enable()
console.print(Theme.dim(" Gradient checkpointing enabled"))
except Exception:
pass
if not train_cfg.get("use_8bit", False):
if not self.is_multi_gpu:
self.model.to(self.device)
else:
console.print(
Theme.dim("Multi-GPU: Model will be moved by Trainer")
)
console.print(Theme.success(" Model loaded"))
# ------------------------------------------------------------------
# Sanity check: apakah full FT di GPU ini realistis?
# ------------------------------------------------------------------
if not will_use_peft and self.device == "cuda":
try:
model_params = sum(p.numel() for p in self.model.parameters())
model_gb = model_params * 4 / (1024 ** 3)
# fp32 params + grads + 2x Adam states ≈ 4x model
estimated_total_gb = model_gb * 4
vram_gb = (
torch.cuda.get_device_properties(0).total_memory
/ (1024 ** 3)
)
console.print(Theme.dim(
f" Estimasi VRAM full FT: ~{estimated_total_gb:.1f}GB "
f"/ {vram_gb:.1f}GB tersedia"
))
if estimated_total_gb > vram_gb * 0.9:
console.print(Theme.error(
" Full fine-tuning kemungkinan besar akan OOM!\n"
" Solusi: aktifkan use_peft=True ATAU "
"use_8bit=True ATAU pakai model lebih kecil."
))
except Exception as e:
debug_logger.debug(f"VRAM check error: {e}")
except Exception as e:
console.print(Theme.error(f"Error loading model: {e}"))
error_logger.error(f"Model loading error: {e}")
return
self._auto_tune_batch_size(train_cfg)
fmt = detect_format(self.tokenizer)
fmt_label = "Native chat template (Qwen)" if fmt == "native" else "Simple (User:/AI:)"
console.print(Theme.dim(f" Format detected: {fmt_label}"))
console.print(Theme.info("Creating dataset..."))
dataset = Dataset.from_dict({"text": samples})
console.print(Theme.info("Tokenizing dataset..."))
num_proc = 1
if os.name == "nt":
num_proc = 1
console.print(Theme.dim("Windows: Using single process for tokenization"))
from functools import partial
tokenize_fn = partial(
_tokenize_batch,
tokenizer=self.tokenizer,
max_length=train_cfg["max_length"],
)
tokenized_dataset = dataset.map(
tokenize_fn,
batched=True,
num_proc=num_proc,
remove_columns=dataset.column_names,
desc="Tokenizing",
)
console.print(Theme.dim(f" Tokenized with {num_proc} processes"))
min_len = 5
max_len = train_cfg["max_length"]
filter_fn = partial(
_filter_valid_lengths, min_len=min_len, max_len=max_len
)
tokenized_dataset = tokenized_dataset.filter(
filter_fn,
batched=True,
num_proc=num_proc if os.name != "nt" else 1,
desc="Filtering valid samples",
)
if len(tokenized_dataset) == 0:
console.print(Theme.error(" Tidak ada sample yang valid! Periksa dataset."))
return
truncated_count = sum(
1
for item in tokenized_dataset
if len(item["input_ids"]) >= train_cfg["max_length"]
)
if truncated_count > 0:
truncation_pct = (truncated_count / len(tokenized_dataset)) * 100
console.print(Theme.warning(f" {truncation_pct:.1f}% samples terpotong"))
max_samples_limit = train_cfg.get("max_samples_limit", 999999999999999999)
if len(tokenized_dataset) > max_samples_limit:
tokenized_dataset = tokenized_dataset.shuffle(
seed=train_cfg.get("seed", 42)
)
tokenized_dataset = tokenized_dataset.select(range(max_samples_limit))
console.print(
Theme.warning(f" Dataset dibatasi {max_samples_limit} sampel (acak)")
)
val_split = train_cfg.get("validation_split", 0.1)
if val_split <= 0 or len(tokenized_dataset) < 20:
if len(tokenized_dataset) >= 20:
val_split = 0.1
else:
val_split = (
max(0.05, 1.0 / len(tokenized_dataset))
if len(tokenized_dataset) > 1
else 0.0
)
console.print(
Theme.warning(f" Dataset kecil, validation split: {val_split:.2f}")
)
if val_split > 0 and len(tokenized_dataset) > 1:
split_data = tokenized_dataset.train_test_split(
test_size=val_split, seed=train_cfg.get("seed", 42)
)
train_dataset = split_data["train"]
eval_dataset = split_data["test"]
else:
train_dataset = tokenized_dataset
eval_dataset = None
console.print(Theme.warning(" No validation set (dataset terlalu kecil)"))
console.print(
Theme.success(
f" Train: {len(train_dataset)}, Eval: {len(eval_dataset) if eval_dataset else 0}"
)
)
# ------------------------------------------------------------------
# Custom data collator: handle labels dari prompt masking
# ------------------------------------------------------------------
def _chat_data_collator(features):
import torch as _torch
max_len = max(len(f["input_ids"]) for f in features)
pad_id = (
self.tokenizer.pad_token_id
or self.tokenizer.eos_token_id
or 0
)
input_ids, attention, labels = [], [], []
for f in features:
pad_len = max_len - len(f["input_ids"])
input_ids.append(list(f["input_ids"]) + [pad_id] * pad_len)
attention.append(list(f["attention_mask"]) + [0] * pad_len)
labels.append(list(f["labels"]) + [-100] * pad_len)
return {
"input_ids": _torch.tensor(input_ids, dtype=_torch.long),
"attention_mask": _torch.tensor(attention, dtype=_torch.long),
"labels": _torch.tensor(labels, dtype=_torch.long),
}
data_collator = _chat_data_collator
checkpoint_dir = os.path.join(CHECKPOINT_DIR, f"{model_name}_checkpoints")
if train_cfg.get("freeze_embeddings", False):
try:
emb = self.model.get_input_embeddings()
for p in emb.parameters():
p.requires_grad = False
console.print(Theme.dim(" Embeddings frozen"))
except Exception as e:
console.print(Theme.warning(f"Could not freeze embeddings: {e}"))
if train_cfg.get("use_peft", False) and HAS_PEFT:
self._apply_peft(train_cfg)
training_args = make_training_args_safe(
checkpoint_dir, train_cfg, self.device, multi_gpu=self.is_multi_gpu
)
callbacks = [
EnhancedCallback(history=self.history, memory_watchdog=self.watchdog),
]
if train_cfg.get("dynamic_grad_accumulation", False):
callbacks.append(
GradientAccumulationScheduler(
initial_steps=train_cfg.get("gradient_accumulation_steps", 1),
max_steps=16,
min_steps=1,
memory_threshold=0.7,
adjustment_interval=100,
)
)
early_stopping_patience = train_cfg.get("early_stopping_patience", 0)
if early_stopping_patience > 0 and eval_dataset and len(eval_dataset) > 5:
try:
callbacks.append(
SafeEarlyStoppingCallback(
early_stopping_patience=early_stopping_patience,
early_stopping_threshold=0.001,
)
)
console.print(Theme.dim(" Early stopping enabled"))
except Exception as e:
console.print(Theme.warning(f" Early stopping setup failed: {e}"))
elif early_stopping_patience > 0:
console.print(
Theme.warning(" Early stopping disabled: eval dataset terlalu kecil")
)
trainer = Trainer(
model=self.model,
args=training_args,
train_dataset=train_dataset,
eval_dataset=eval_dataset,
data_collator=data_collator,
callbacks=callbacks,
)
console.print(Theme.success(" Memulai training..."))
train_logger.info(f"Starting training: {model_name}")
try:
train_result = trainer.train()
training_time = time.time() - self._start_time
console.print(Theme.success(" Training selesai!"))
final_train_loss = train_result.training_loss
if eval_dataset and len(eval_dataset) > 0:
try:
eval_metrics = trainer.evaluate()
final_eval_loss = eval_metrics.get("eval_loss", final_train_loss)
except Exception as e:
console.print(Theme.warning(f" Evaluation failed: {e}"))
eval_metrics = {}
final_eval_loss = final_train_loss
else:
eval_metrics = {}
final_eval_loss = final_train_loss
if (
final_eval_loss is not None
and final_eval_loss < 1000
and final_eval_loss > 0
):
try:
perplexity = math.exp(final_eval_loss)
except OverflowError:
perplexity = float("inf")
else:
perplexity = None
metrics = {
"final_train_loss": float(final_train_loss),
"final_eval_loss": float(final_eval_loss),
"perplexity": float(perplexity)
if perplexity is not None and math.isfinite(perplexity)
else None,
"training_time_seconds": float(training_time),
"steps_trained": int(trainer.state.global_step)
if hasattr(trainer.state, "global_step")
else 0,
}
if eval_dataset and len(eval_dataset) > 5:
gen_metrics = self._eval_generation(trainer, eval_dataset, train_cfg)
metrics.update(gen_metrics)
else:
console.print(
Theme.dim("Skipping generation evaluation (dataset terlalu kecil)")
)
console.print(Theme.info("Menyimpan model final..."))
if self.is_peft_model and HAS_PEFT:
try:
self.model.save_pretrained(model_path)
self.tokenizer.save_pretrained(model_path)
with open(
os.path.join(model_path, "base_model_info.json"), "w"
) as f:
json.dump({"base_model": base_model, "use_peft": True}, f)
console.print(Theme.success(" PEFT adapter saved"))
except Exception as e:
console.print(
Theme.warning(
f"Failed to save PEFT adapter: {e}. Saving full model."
)
)
trainer.save_model(model_path)
self.tokenizer.save_pretrained(model_path)
else:
trainer.save_model(model_path)
self.tokenizer.save_pretrained(model_path)
history_path = os.path.join(model_path, "training_history.json")
with open(history_path, "w") as f:
json.dump(self.history.to_dict(), f, indent=2)
console.print(Theme.success(f" Model tersimpan di {model_path}"))
metadata = ModelMetadata(
model_name=model_name,
base_model=base_model,
dataset_path=dataset_path,
dataset_format=stats.format_detected,
dataset_size=stats.total_samples,
training_samples=len(train_dataset),
validation_samples=len(eval_dataset) if eval_dataset else 0,
created_date=time.strftime("%Y-%m-%d %H:%M:%S"),
description=f"Fine-tuned {base_model}",
metrics=metrics,
config=train_cfg,
tags=["ultimate", "enhanced", "qwen"]
if "qwen" in base_model.lower()
else ["ultimate", "enhanced"],
)
self.model_manager.add_model(model_name, metadata)
self.model_manager.set_current_model(model_name)
console.print(Theme.success(f" Model '{model_name}' sekarang aktif"))
self._show_training_summary(metrics)
self.history.print_summary()
except KeyboardInterrupt:
console.print(Theme.warning(" Training diinterupsi!"))
self._save_emergency_checkpoint(trainer, model_path)
except Exception as e:
console.print(Theme.error(f"Error: {rich_escape(str(e))}"))
error_logger.error(f"Training error: {e}")
self._save_emergency_checkpoint(trainer, model_path)
raise
finally:
if self.is_multi_gpu and dist.is_initialized():
dist.destroy_process_group()
self._cleanup()
def _eval_generation(
self, trainer: Trainer, eval_dataset: Dataset, train_cfg: Dict
) -> Dict[str, Any]:
results = {}
try:
n = min(len(eval_dataset), train_cfg.get("eval_gen_samples", 16))
if n == 0:
return results
generated_texts = []
reference_texts = []
for i in range(n):
item = eval_dataset[i]
text = self.tokenizer.decode(
item["input_ids"], skip_special_tokens=True
)
if "AI:" in text:
ai_marker = text.find("AI:")
user_part = text[:ai_marker].strip()
ref_part = text[ai_marker + 3:].strip()
elif "<|im_start|>assistant" in text:
parts = text.split("<|im_start|>assistant")
user_part = parts[0].strip()
ref_part = parts[1].split("<|im_end|>")[0].strip() if len(parts) > 1 else ""
else:
enc = self.tokenizer.encode(text, truncation=True, max_length=128)
if not enc:
continue
user_part = text
ref_part = ""
if not user_part:
continue
input_enc = self.tokenizer.encode(user_part, return_tensors="pt").to(self.device)
if ref_part:
reference_texts.append(ref_part)
max_new = min(50, train_cfg.get("max_length", 768) - input_enc.shape[1])
if max_new <= 0:
continue
with torch.no_grad():
out = trainer.model.generate(
input_enc,
max_new_tokens=max_new,
do_sample=True,
temperature=0.7,
pad_token_id=self.tokenizer.eos_token_id,
eos_token_id=self.tokenizer.eos_token_id,
)
gen = self.tokenizer.decode(
out[0, input_enc.shape[1]:], skip_special_tokens=True
).strip()
if gen:
generated_texts.append(gen)
if generated_texts and reference_texts:
results.update(compute_bleu_rouge(generated_texts, reference_texts))
results["samples_evaluated"] = len(generated_texts)
except Exception as e:
error_logger.error(f"Eval generation error: {e}")
return results
def _resume_training(self) -> None:
from rich.prompt import Prompt
from pathlib import Path
checkpoint_dir = Path(CHECKPOINT_DIR)
checkpoints = list(checkpoint_dir.glob("*_checkpoints/checkpoint-*"))
if not checkpoints:
console.print(Theme.error("Tidak ada checkpoint ditemukan!"))
return
console.print(Theme.info("Checkpoint tersedia:"))
for idx, cp in enumerate(sorted(checkpoints), 1):
console.print(f" {idx}. {cp}")
choice = Prompt.ask("[yellow]Pilih checkpoint nomor", default="1")
try:
checkpoint = sorted(checkpoints)[int(choice) - 1]
except (ValueError, IndexError):
console.print(Theme.error("Pilihan invalid"))
return
dataset_path = self._select_dataset_file()
if dataset_path is None:
return
if not os.path.exists(dataset_path):
console.print(Theme.error("Dataset tidak ditemukan!"))
return
console.print(Theme.info("Membaca dataset..."))
try:
loader = EnhancedDatasetLoader()
train_cfg = self.model_manager.get_training_config()
samples, stats = loader.load(dataset_path)
except Exception as e:
console.print(Theme.error(f"Error: {rich_escape(str(e))}"))
return
self._show_dataset_stats(stats)
model_name = Prompt.ask("[cyan]Nama model output", default="my-qwen-resumed")
train_cfg = self.model_manager.get_training_config()
train_cfg["resume_from_checkpoint"] = str(checkpoint)
self._execute_training(
samples=samples,
model_name=model_name,
dataset_path=dataset_path,
train_cfg=train_cfg,
stats=stats,
)
def _save_emergency_checkpoint(self, trainer: Trainer, model_path: str) -> None:
try:
console.print(Theme.warning(" Saving emergency checkpoint..."))
if self.is_peft_model and HAS_PEFT:
try:
self.model.save_pretrained(model_path)
except Exception:
trainer.save_model(model_path)
else:
trainer.save_model(model_path)
if self.tokenizer:
self.tokenizer.save_pretrained(model_path)
console.print(Theme.success(" Emergency checkpoint saved"))
except Exception as e:
console.print(Theme.error(f"Emergency save failed: {e}"))
def _show_dataset_stats(self, stats: DatasetStats) -> None:
from rich.table import Table
from rich import box
console.print("\n" + Theme.header(" Dataset Statistics:"))
table = Table(title="Dataset Info", box=box.ROUNDED)
table.add_column("Metric", style="cyan")
table.add_column("Value", style="green")
table.add_row("Format", stats.format_detected)
table.add_row("Structure", stats.structure_type)
table.add_row("Total Samples", f"{stats.total_samples:,}")
table.add_row("Valid Samples", f"{stats.valid_samples:,}")
table.add_row("Invalid Samples", f"{stats.invalid_samples:,}")
table.add_row("Duplicate Samples", f"{stats.duplicate_samples:,}")
table.add_row("Avg Length", f"{stats.avg_length:.1f} chars")
table.add_row("Avg Words", f"{stats.avg_words:.1f}")
table.add_row("Std Length", f"{stats.std_length:.1f}")
table.add_row("Min Length", f"{stats.min_length}")
table.add_row("Max Length", f"{stats.max_length}")
table.add_row("Language", stats.language)
if stats.conversation_pairs:
table.add_row("Conversation Pairs", str(len(stats.conversation_pairs)))
console.print(table)
if stats.warnings:
console.print("\n" + Theme.warning("Warnings:"))
for warn in stats.warnings[:5]:
console.print(f" {warn}")
def _show_training_summary(self, metrics: Dict) -> None:
from rich.table import Table
from rich import box
console.print("\n" + Theme.header(" Training Metrics Summary:"))
table = Table(title="Training Results", box=box.ROUNDED)
table.add_column("Metric", style="cyan")
table.add_column("Value", style="green")
table.add_row("Train Loss", f"{metrics.get('final_train_loss', 0):.4f}")
table.add_row("Eval Loss", f"{metrics.get('final_eval_loss', 0):.4f}")
perplexity = metrics.get("perplexity")
if perplexity is not None and math.isfinite(perplexity):
table.add_row("Perplexity", f"{perplexity:.2f}")
if metrics.get("unigram_overlap"):
table.add_row("Unigram Overlap", f"{metrics['unigram_overlap']:.4f}")
if metrics.get("bleu"):
table.add_row("BLEU", f"{metrics['bleu']:.4f}")
if metrics.get("rougeL"):
table.add_row("ROUGE-L", f"{metrics['rougeL']:.4f}")
training_time = metrics.get("training_time_seconds", 0)
table.add_row("Training Time", f"{training_time / 60:.1f} minutes")
table.add_row("Steps", f"{metrics.get('steps_trained', 0):,}")
console.print(table)
def _cleanup(self) -> None:
if self.watchdog:
self.watchdog.stop()
if torch.cuda.is_available():
torch.cuda.empty_cache()
gc.collect()
console.print(Theme.dim(" Cleanup complete"))
def _select_dataset_file(self, prompt: str = "Pilih file dataset") -> Optional[str]:
from pathlib import Path
from rich.prompt import Prompt
from config import DATA_DIR
supported_ext = (
".txt", ".json", ".jsonl", ".csv", ".tsv",
".parquet", ".arrow", ".json.gz", ".jsonl.gz",
)
data_dir = Path(DATA_DIR)
files = []
for ext in supported_ext:
files.extend(data_dir.glob(f"*{ext}"))
files = sorted(files)
if not files:
console.print(
Theme.warning(
"Tidak ada file dataset di folder 'data/'. Silakan masukkan path manual."
)
)
path = Prompt.ask("[cyan]Masukkan path file dataset")
if os.path.exists(path):
return path
console.print(Theme.error("File tidak ditemukan!"))
return None
console.print(Theme.info("File dataset tersedia:"))
for idx, f in enumerate(files, 1):
try:
size = f.stat().st_size
if size > 100 * 1024 * 1024:
size_str = f"{size / (1024 * 1024 * 1024):.2f} GB"
elif size > 1024 * 1024:
size_str = f"{size / (1024 * 1024):.2f} MB"
elif size > 1024:
size_str = f"{size / 1024:.2f} KB"
else:
size_str = f"{size} B"
console.print(f" {idx}. {f.name} ({size_str})")
except Exception:
console.print(f" {idx}. {f.name}")
console.print(" 0. Masukkan path manual")
choice = Prompt.ask("[yellow]Pilih nomor", default="1")
if choice == "0":
path = Prompt.ask("[cyan]Masukkan path file dataset")
if os.path.exists(path):
return path
console.print(Theme.error("File tidak ditemukan!"))
return None
try:
idx = int(choice) - 1
if 0 <= idx < len(files):
return str(files[idx])
console.print(Theme.error("Nomor tidak valid!"))
return None
except ValueError:
console.print(Theme.error("Input harus berupa angka!"))
return None