CoolFace
Apppublic

google/functiongemma-tuning-lab

sourceHugging Faceapache-2.0updated 8mo agoView on Hugging Face
86likes
engine.py450 linesDownload Raw Back to root
1import threading2import torch3import time4import json5import queue6import uuid7import matplotlib.pyplot as plt8from functools import partial9from typing import Generator, Optional, List, Dict, Any, Tuple10from datasets import Dataset, load_dataset11from trl import SFTConfig, SFTTrainer12from transformers import TrainerCallback, TrainingArguments, TrainerState, TrainerControl13from huggingface_hub import HfApi, model_info, metadata_update14 15from config import AppConfig16from tools import DEFAULT_TOOLS17from utils import (18    authenticate_hf, 19    load_model_and_tokenizer, 20    create_conversation_format, 21    parse_csv_dataset,22    zip_directory23)24 25class AbortCallback(TrainerCallback):26    def __init__(self, stop_event: threading.Event):27        self.stop_event = stop_event28 29    def on_step_end(self, args: TrainingArguments, state: TrainerState, control: TrainerControl, **kwargs):30        if self.stop_event.is_set():31            control.should_training_stop = True32 33class LogStreamingCallback(TrainerCallback):34    def __init__(self, log_queue: queue.Queue):35        self.log_queue = log_queue36        37    def _get_string(self, value):38        if isinstance(value, float):39            return f"{value:.4f}"40        return str(value)41 42    def on_log(self, args, state, control, logs=None, **kwargs):43        if not logs:44            return45 46        metrics_map = {47            "loss": "Loss",48            "eval_loss": "Eval Loss",49            "learning_rate": "LR",50            "epoch": "Epoch"51        }52        log_parts = [f"๐Ÿ“ [Step {state.global_step}]"]53        54        for key, label in metrics_map.items():55            if key in logs:56                val = logs[key]57                if isinstance(val, (float, int)):58                    val_str = f"{val:.4f}" if val > 1e-4 else f"{val:.2e}"59                else:60                    val_str = str(val)61            62                log_parts.append(f"{label}: {val_str}")63        64        log_payload = logs.copy()65        log_payload['step'] = state.global_step66        67        self.log_queue.put((" | ".join(log_parts), log_payload))68 69class FunctionGemmaEngine:70    def __init__(self, config: AppConfig):71        self.config = config72        73        self.session_id = str(uuid.uuid4())[:8]74        self.output_dir = self.config.ARTIFACTS_DIR.joinpath(f"session_{self.session_id}")75        self.output_dir.mkdir(parents=True, exist_ok=True)76 77        self.model = None78        self.tokenizer = None79        self.loaded_model_name = None 80        self.imported_dataset = []81        self.stop_event = threading.Event()82        self.current_tools = DEFAULT_TOOLS83        self.has_model_tuned = False84 85        authenticate_hf(self.config.HF_TOKEN)86        try:87            self.refresh_model()88        except Exception as e:89            print(f"Initial load warning: {e}")90 91    # --- Tool Schema Methods ---92    def get_tools_json(self) -> str:93        return json.dumps(self.current_tools, indent=2)94 95    def update_tools(self, json_str: str) -> str:96        try:97            new_tools = json.loads(json_str)98            if not isinstance(new_tools, list):99                return "Error: Schema must be a list of tool definitions."100            self.current_tools = new_tools101            return "โœ… Tool Schema Updated successfully."102        except json.JSONDecodeError as e:103            return f"โŒ JSON Error: {e}"104        except Exception as e:105            return f"โŒ Error: {e}"106 107    # --- Model & Data Management ---108    109    def _load_model_weights(self):110        print(f"[{self.session_id}] Loading model: {self.config.MODEL_NAME}...")111        self.model, self.tokenizer = load_model_and_tokenizer(self.config.MODEL_NAME)112        self.loaded_model_name = self.config.MODEL_NAME113 114    def refresh_model(self) -> str:115        self.has_model_tuned = False116        try:117            self._load_model_weights()118            return f"Model loaded: {self.loaded_model_name}\nData cleared.\nReady (Session {self.session_id})."119        except Exception as e:120            self.model = None121            self.tokenizer = None122            self.loaded_model_name = None123            return f"CRITICAL ERROR: Model failed to load. {e}"124 125    def load_csv(self, file_path: str) -> str:126        try:127            new_data = parse_csv_dataset(file_path)128            if not new_data:129                return "Error: File empty or format invalid."130            self.imported_dataset = new_data131            return f"Successfully imported {len(new_data)} samples."132        except Exception as e:133            return f"Import failed: {e}"134 135    def trigger_stop(self):136        self.stop_event.set()137 138    def _ensure_model_consistency(self) -> Generator[str, None, bool]:139        """Checks if the requested model matches the loaded one. Reloads if necessary."""140        if self.config.MODEL_NAME != self.loaded_model_name:141            yield f"๐Ÿ”„ Model changed. Switching from '{self.loaded_model_name}' to '{self.config.MODEL_NAME}'...\n"142            try:143                self._load_model_weights()144                yield "โœ… Model reloaded successfully.\n"145                return True146            except Exception as e:147                yield f"โŒ Failed to load model '{self.config.MODEL_NAME}': {e}\n"148                return False149        if self.model is None:150             yield "โŒ Error: No model loaded.\n"151             return False152        return True153 154    # --- Evaluation Pipeline ---155    156    def run_evaluation(self, test_size: float, shuffle_data: bool) -> Generator[str, None, None]:157        self.stop_event.clear()158        output_buffer = ""159        160        try:161            # 1. Check Model162            gen = self._ensure_model_consistency()163            try:164                while True:165                    msg = next(gen)166                    output_buffer += msg167                    yield output_buffer168            except StopIteration as e:169                if not e.value: return # Failed to load170                171            # 2. Prepare Data172            output_buffer += f"โณ Preparing Dataset for Eval (Test Split: {test_size})...\n"173            yield output_buffer174 175            dataset, log = self._prepare_dataset()176            output_buffer += log177            yield output_buffer178                179            if not dataset:180                output_buffer += "โŒ Dataset creation failed.\n"181                yield output_buffer182                return183 184            if len(dataset) > 1:185                dataset = dataset.train_test_split(test_size=test_size, shuffle=shuffle_data)186            else:187                dataset = {"train": dataset, "test": dataset}188                189            # 3. Run Inference190            output_buffer += "\n๐Ÿ“Š Evaluating Model Success Rate on Test Split...\n"191            yield output_buffer192 193            for update in self._evaluate_model(dataset["test"]):194                yield f"{output_buffer}{update}"195                if self.stop_event.is_set():196                    yield f"{output_buffer}{update}\n\n๐Ÿ›‘ Evaluation interrupted by user."197                    break198        finally:199            self.stop_event.set() # Ensure loop breaks if generator cancelled200 201    # --- Training Pipeline ---202 203    def run_training_pipeline(self, epochs: int, learning_rate: float, test_size: float, shuffle_data: bool) -> Generator[Tuple[str, Any], None, None]:204        self.stop_event.clear()205        output_buffer = ""206        last_plot = None207 208        try:209            # 1. Check Model210            gen = self._ensure_model_consistency()211            try:212                while True:213                    msg = next(gen)214                    output_buffer += f"{msg}"215                    yield output_buffer, None216            except StopIteration as e:217                if not e.value: return218 219            output_buffer += f"โณ Preparing Dataset (Test Split: {test_size}, Shuffle: {shuffle_data})...\n"220            yield output_buffer, None221 222            dataset, log = self._prepare_dataset()223            if not dataset:224                yield "Dataset creation failed.", None225                return226 227            output_buffer += log228            yield output_buffer, None229                230            if len(dataset) > 1:231                dataset = dataset.train_test_split(test_size=test_size, shuffle=shuffle_data)232            else:233                dataset = {"train": dataset, "test": dataset}234 235            # --- Training (Threaded) ---236            output_buffer += f"\n๐Ÿš€ Starting Fine-tuning (Epochs: {epochs}, LR: {learning_rate})...\n"237            yield output_buffer, None238            239            log_queue = queue.Queue()240            training_error = None241            running_history = [] 242            243            def train_wrapper():244                nonlocal training_error245                try:246                    self._execute_trainer(dataset, log_queue, epochs, learning_rate)247                except Exception as e:248                    training_error = e249                    250            train_thread = threading.Thread(target=train_wrapper)251            train_thread.start()252            253            while train_thread.is_alive():254                while not log_queue.empty():255                    payload = log_queue.get()256                    if isinstance(payload, tuple):257                        msg, log_data = payload258                        output_buffer += f"{msg}\n"259                        running_history.append(log_data)260                        try:261                            last_plot = self._generate_loss_plot(running_history)262                            yield output_buffer, last_plot263                        except Exception:264                            yield output_buffer, last_plot265                    else:266                        output_buffer += f"{payload}\n"267                        yield output_buffer, last_plot268                269                if self.stop_event.is_set():270                    yield f"{output_buffer}๐Ÿ›‘ Stop signal sent. Waiting for trainer to wrap up...\n", last_plot271                272                time.sleep(0.1)273            274            train_thread.join()275            276            self.has_model_tuned = True277            278            while not log_queue.empty():279                payload = log_queue.get()280                if isinstance(payload, tuple):281                    msg, log_data = payload282                    output_buffer += f"{msg}\n"283                    running_history.append(log_data)284                    last_plot = self._generate_loss_plot(running_history)285                else:286                    output_buffer += f"{payload}\n"287                yield output_buffer, last_plot288                    289            if training_error:290                output_buffer += f"โŒ Error during training: {training_error}\n"291                yield output_buffer, last_plot292                return293 294            if self.stop_event.is_set(): 295                output_buffer += "๐Ÿ›‘ Training manually stopped.\n"296                yield output_buffer, last_plot297                return298            299            output_buffer += "โœ… Training finished.\n"300            yield output_buffer, last_plot301            302        finally:303            self.stop_event.set() # Ensure background thread stops if generator cancelled304 305    def _prepare_dataset(self):306        formatting_fn = partial(create_conversation_format, tools_list=self.current_tools)307 308        if not self.imported_dataset:309            ds = load_dataset(self.config.DEFAULT_DATASET, split="train").map(formatting_fn)310            log = f" `-> using default dataset (size:{len(ds)})\n"311        else:312            dataset_as_dicts = [{313                "user_content": row[0], "tool_name": row[1], "tool_arguments": row[2]}314                for row in self.imported_dataset315            ]316            ds = Dataset.from_list(dataset_as_dicts).map(formatting_fn)317            log = f" `-> using custom dataset (size:{len(ds)})\n"318        return ds, log319 320    def _execute_trainer(self, dataset, log_queue: queue.Queue, epochs: int, learning_rate: float) -> List[Dict]:321        torch_dtype = self.model.dtype322        args = SFTConfig(323            output_dir=str(self.output_dir),324            max_length=512,325            packing=False,326            num_train_epochs=epochs,327            per_device_train_batch_size=4,328            logging_steps=1,329            save_strategy="no",330            eval_strategy="epoch",331            learning_rate=learning_rate,332            fp16=(torch_dtype == torch.float16),333            bf16=(torch_dtype == torch.bfloat16),334            report_to="none",335            dataset_kwargs={"add_special_tokens": False, "append_concat_token": True}336        )337 338        trainer = SFTTrainer(339            model=self.model,340            args=args,341            train_dataset=dataset['train'],342            eval_dataset=dataset['test'],343            processing_class=self.tokenizer,344            callbacks=[345                AbortCallback(self.stop_event),346                LogStreamingCallback(log_queue)347            ]348        )349        trainer.train()350        trainer.save_model()351        return trainer.state.log_history352        353    def _generate_loss_plot(self, history: list):354        if not history: return None355        plt.close('all')356        357        train_steps = [x['step'] for x in history if 'loss' in x]358        train_loss = [x['loss'] for x in history if 'loss' in x]359        eval_steps = [x['step'] for x in history if 'eval_loss' in x]360        eval_loss = [x['eval_loss'] for x in history if 'eval_loss' in x]361 362        fig, ax = plt.subplots(figsize=(10, 5))363        if train_steps:364            ax.plot(train_steps, train_loss, label='Training Loss', linestyle='-', marker=None)365        if eval_steps:366            ax.plot(eval_steps, eval_loss, label='Validation Loss', linestyle='--', marker='o')367 368        ax.set_xlabel("Steps")369        ax.set_ylabel("Loss")370        ax.set_title("Training & Validation Loss")371        ax.legend()372        ax.grid(True, linestyle=':', alpha=0.6)373        plt.tight_layout()374        return fig375 376    def _evaluate_model(self, test_dataset) -> Generator[str, None, None]:377        results = []378        success_count = 0379        for idx, item in enumerate(test_dataset):380            messages = item["messages"][:2]381            try:382                inputs = self.tokenizer.apply_chat_template(383                    messages, tools=self.current_tools, add_generation_prompt=True, return_dict=True, return_tensors="pt"384                )385                device = self.model.device386                inputs = {k: v.to(device) for k, v in inputs.items()}387                out = self.model.generate(388                    **inputs, 389                    pad_token_id=self.tokenizer.eos_token_id, 390                    max_new_tokens=128391                )392                output = self.tokenizer.decode(out[0][len(inputs["input_ids"][0]):], skip_special_tokens=True)393                log_entry = f"{idx+1}. Prompt: {messages[1]['content']}\n   Output: {output[:100]}..."394                expected_tool = item['messages'][2]['tool_calls'][0]['function']['name']395                if expected_tool in output:396                    log_entry += "\n   -> โœ… Correct Tool"397                    success_count += 1398                else:399                    log_entry += f"\n   -> โŒ Wrong Tool (Expected: {expected_tool})"400                results.append(log_entry)401                yield "\n".join(results) + f"\n\nRunning Success Rate: {success_count}/{idx+1}"402            except Exception as e:403                yield f"Error during inference: {e}"404 405    def get_zip_path(self) -> Optional[str]:406        if not self.output_dir.exists(): return None407        base_name = str(self.config.ARTIFACTS_DIR.joinpath(f"functiongemma_finetuned_{self.session_id}"))408        return zip_directory(str(self.output_dir), base_name)409 410    def upload_model_to_hub(self, repo_name: str, oauth_token: str) -> str:411        """Uploads the trained model to Hugging Face Hub."""412        if not self.output_dir.exists() or not any(self.output_dir.iterdir()):413            return "โŒ No trained model found in current session. Run training first."414        415        try:416            api = HfApi(token=oauth_token)417 418            # Get the authenticated user's username419            user_info = api.whoami()420            username = user_info['name']421        422            # Construct the full repo ID423            repo_id = f"{username}/{repo_name}"424            print(f"Preparing to upload to: {repo_id}")425 426            # Create the repo (safe if it already exists)427            api.create_repo(repo_id=repo_id, exist_ok=True)428            429            # Upload430            print(f"Uploading to {repo_id}...")431            repo_url = api.upload_folder(432                folder_path=str(self.output_dir),433                repo_id=repo_id,434                repo_type="model"435            )436 437            info = model_info(438                repo_id=repo_id,439                token=oauth_token440            )441            tags = ["functiongemma", "functiongemma-tuning-lab"]442            if info.card_data:443                tags = info.card_data.tags444                tags.append("functiongemma-tuning-lab")445 446            metadata_update(repo_id, {"tags": tags}, overwrite=True, token=oauth_token)447 448            return f"โœ… Success! Model uploaded to: {repo_url}"449        except Exception as e:450            return f"โŒ Upload failed: {str(e)}"