google/functiongemma-tuning-lab
86
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)}"