Download trainer.py from FannyFa/FannyFa-Model-V1: direct link, hf CLI and curl.
- Browser
- Download file 64.5 kB
-
https://huggingface.co/FannyFa/FannyFa-Model-V1/resolve/main/trainer.py
- Command line
-
hf download hf://FannyFa/FannyFa-Model-V1/trainer.py
-
curl -L -o trainer.py https://huggingface.co/FannyFa/FannyFa-Model-V1/resolve/main/trainer.py
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 |