{ "nbformat": 4, "nbformat_minor": 0, "metadata": { "colab": { "provenance": [], "gpuType": "A100" }, "kernelspec": { "name": "python3", "display_name": "Python 3" }, "language_info": { "name": "python" }, "accelerator": "GPU" }, "cells": [ { "cell_type": "code", "metadata": { "id": "09f76c4a" }, "source": [ "# Install uv\n", "!curl -LsSf https://astral.sh/uv/install.sh | sh\n", "import os\n", "os.environ['PATH'] = f\"{os.path.expanduser('~/.local/bin')}:{os.environ['PATH']}\"" ], "execution_count": null, "outputs": [] }, { "cell_type": "code", "metadata": { "id": "78b92072" }, "source": [ "pyproject_content = \"\"\"[project]\n", "name = \"finetune\"\n", "version = \"0.1.0\"\n", "description = \"Add your description here\"\n", "readme = \"README.md\"\n", "requires-python = \">=3.10\"\n", "dependencies = [\n", " \"datasets>=4.3.0\",\n", " \"dill>=0.4.0\",\n", " \"unsloth>=2026.5.5\",\n", " \"unsloth-zoo>=2026.5.3\",\n", "]\n", "\"\"\"\n", "\n", "with open(\"pyproject.toml\", \"w\") as f:\n", " f.write(pyproject_content)\n", "\n", "with open(\".python-version\", \"w\") as f:\n", " f.write(\"3.10\")\n", "\n", "with open(\"README.md\", \"w\") as f:\n", " f.write(\"# Finetune Project\")" ], "execution_count": null, "outputs": [] }, { "cell_type": "code", "metadata": { "collapsed": true, "id": "b683e0d1" }, "source": [ "# Sync the environment and install extra unsloth requirements\n", "!uv sync" ], "execution_count": null, "outputs": [] }, { "cell_type": "code", "source": [ "script_content = \"\"\"\n", "COMPARISON_TRAIN = True # False → causal_text, True → simple_text\n", "\n", "import json\n", "import os\n", "from dataclasses import dataclass\n", "from datasets import load_dataset\n", "from unsloth import FastVisionModel\n", "from trl import SFTTrainer, SFTConfig\n", "import torch\n", "\n", "# Determine output names based on switch\n", "suffix = \"_comparison\" if COMPARISON_TRAIN else \"\"\n", "output_dir = f\"outputs_qwen35_0.8b_cocreator{suffix}\"\n", "lora_dir = f\"lora_cocreator_qwen35_0.8b{suffix}\"\n", "safetensors_dir = f\"safetensors_cocreator_qwen35_0.8b{suffix}\"\n", "\n", "text_col = \"simple_text\" if COMPARISON_TRAIN else \"causal_text\"\n", "print(f\"Training with text column: {text_col}\")\n", "print(f\"Output: {safetensors_dir}\")\n", "\n", "max_seq_length = 2048\n", "\n", "# 1. Load dataset from Hugging Face\n", "dataset = load_dataset(\n", " \"NIyueeE/cocreator-driving-scene\",\n", " split=\"train\",\n", ")\n", "print(f\"Total samples: {len(dataset)}\")\n", "\n", "# 2. Build messages (system_prompt + selected text column)\n", "system_prompts = dataset[\"system_prompt\"]\n", "texts = dataset[text_col]\n", "messages_list = [\n", " [\n", " {\"role\": \"system\", \"content\": [{\"type\": \"text\", \"text\": sys_prompt}]},\n", " {\"role\": \"user\", \"content\": [{\"type\": \"image\"}]},\n", " {\"role\": \"assistant\", \"content\": [{\"type\": \"text\", \"text\": text}]},\n", " ]\n", " for sys_prompt, text in zip(system_prompts, texts)\n", "]\n", "\n", "dataset = dataset.rename_column(\"video_frames\", \"images\")\n", "dataset = dataset.add_column(\"messages\", messages_list)\n", "dataset = dataset.remove_columns([\"id\", \"causal_text\", \"system_prompt\", \"simple_text\"])\n", "\n", "# 3. Load model\n", "model, tokenizer = FastVisionModel.from_pretrained(\n", " model_name=\"Qwen/Qwen3.5-0.8B\",\n", " max_seq_length=max_seq_length,\n", " load_in_4bit=False,\n", " full_finetuning=False,\n", ")\n", "\n", "# 4. LoRA\n", "model = FastVisionModel.get_peft_model(\n", " model,\n", " finetune_vision_layers=True,\n", " finetune_language_layers=True,\n", " finetune_attention_modules=True,\n", " finetune_mlp_modules=True,\n", " r=16,\n", " lora_alpha=16,\n", " lora_dropout=0,\n", " bias=\"none\",\n", " random_state=3407,\n", " target_modules=\"all-linear\",\n", " use_gradient_checkpointing=False,\n", " max_seq_length=max_seq_length,\n", ")\n", "\n", "model = model.to(torch.float32)\n", "\n", "@dataclass\n", "class Qwen35VLDataCollator:\n", " processor: callable\n", "\n", " def __call__(self, samples):\n", " sample = samples[0]\n", " imgs = sample[\"images\"]\n", " msgs = sample[\"messages\"]\n", "\n", " for msg in msgs:\n", " if msg[\"role\"] == \"user\":\n", " msg[\"content\"] = [{\"type\": \"image\"} for _ in imgs]\n", "\n", " # Keep original resolution, just ensure RGB format\n", " parsed_imgs = []\n", " for img in imgs:\n", " if img.mode != \"RGB\":\n", " img = img.convert(\"RGB\")\n", " parsed_imgs.append(img)\n", "\n", " text = self.processor.apply_chat_template(msgs, tokenize=False, add_generation_prompt=False)\n", "\n", " result = self.processor(\n", " images=parsed_imgs,\n", " text=text,\n", " padding=True,\n", " return_tensors=\"pt\",\n", " add_special_tokens=False,\n", " )\n", "\n", " if \"pixel_values\" in result:\n", " result[\"pixel_values\"] = result[\"pixel_values\"].to(torch.float32)\n", "\n", " result[\"labels\"] = result[\"input_ids\"].clone()\n", " result[\"labels\"][result[\"attention_mask\"] == 0] = -100\n", " return result\n", "\n", "# 5. Train\n", "trainer = SFTTrainer(\n", " model=model,\n", " train_dataset=dataset,\n", " tokenizer=tokenizer,\n", " data_collator=Qwen35VLDataCollator(tokenizer),\n", " args=SFTConfig(\n", " max_seq_length=max_seq_length,\n", " per_device_train_batch_size=128,\n", " gradient_accumulation_steps=1,\n", " warmup_steps=5,\n", " num_train_epochs=7,\n", " logging_steps=5,\n", " save_strategy=\"epoch\",\n", " output_dir=output_dir,\n", " optim=\"adamw_8bit\",\n", " dataloader_num_workers=4,\n", " seed=3407,\n", " remove_unused_columns=False,\n", " report_to=\"none\",\n", " fp16=False,\n", " bf16=True,\n", " ),\n", ")\n", "\n", "trainer.train()\n", "\n", "model.save_pretrained(lora_dir)\n", "tokenizer.save_pretrained(lora_dir)\n", "model.save_pretrained_merged(safetensors_dir, tokenizer, save_method=\"merged_16bit\")\n", "\n", "# Save the switch config for downstream cells\n", "with open(\"_train_config.json\", \"w\") as f:\n", " json.dump({\"COMPARISON_TRAIN\": COMPARISON_TRAIN}, f)\n", "\"\"\"\n", "\n", "with open(\"train.py\", \"w\") as f:\n", " f.write(script_content)\n", "\n", "!uv run python train.py\n" ], "metadata": { "id": "IMh_Ex3dPx57" }, "execution_count": null, "outputs": [] }, { "cell_type": "markdown", "metadata": { "id": "b329ac18" }, "source": [ "### 1. 获取并可视化训练日志 (Loss vs Step)\n", "Hugging Face Trainer 会在输出目录(如 `checkpoint-xxx`)中保存 `trainer_state.json`。我们可以读取它来提取 `loss` 和 `step`。" ] }, { "cell_type": "code", "metadata": { "id": "4507467e" }, "source": [ "import json\n", "import os\n", "import matplotlib.pyplot as plt\n", "from google.colab import files\n", "\n", "COMPARISON_TRAIN = True # False → causal_text, True → simple_text\n", "\n", "# 找到最新的 checkpoint 目录\n", "suffix = \"_comparison\" if COMPARISON_TRAIN else \"\"\n", "\n", "output_dir = f\"outputs_qwen35_0.8b_cocreator{suffix}\"\n", "checkpoints = [d for d in os.listdir(output_dir) if d.startswith(\"checkpoint-\")]\n", "checkpoints.sort(key=lambda x: int(x.split(\"-\")[1]))\n", "latest_checkpoint = os.path.join(output_dir, checkpoints[-1])\n", "\n", "# 读取 trainer_state.json\n", "state_path = os.path.join(latest_checkpoint, \"trainer_state.json\")\n", "with open(state_path, \"r\") as f:\n", " state = json.load(f)\n", "\n", "# 提取 step, loss, learning rate, grad norm\n", "steps = []\n", "losses = []\n", "lrs = []\n", "grad_norms = []\n", "\n", "for log in state.get(\"log_history\", []):\n", " if \"loss\" in log and \"step\" in log:\n", " steps.append(log[\"step\"])\n", " losses.append(log[\"loss\"])\n", " lrs.append(log.get(\"learning_rate\", 0))\n", " grad_norms.append(log.get(\"grad_norm\", 0))\n", "\n", "# 创建 1行3列 的综合图表\n", "fig, axes = plt.subplots(1, 3, figsize=(18, 5))\n", "\n", "# 1. 绘制 Loss 曲线\n", "axes[0].plot(steps, losses, marker='o', linestyle='-', color='b')\n", "axes[0].set_title(\"Training Loss\")\n", "axes[0].set_xlabel(\"Step\")\n", "axes[0].set_ylabel(\"Loss\")\n", "axes[0].grid(True)\n", "\n", "# 2. 绘制 Learning Rate 曲线\n", "axes[1].plot(steps, lrs, marker='s', linestyle='-', color='orange')\n", "axes[1].set_title(\"Learning Rate Schedule\")\n", "axes[1].set_xlabel(\"Step\")\n", "axes[1].set_ylabel(\"Learning Rate\")\n", "axes[1].grid(True)\n", "\n", "# 3. 绘制 Gradient Norm 曲线\n", "axes[2].plot(steps, grad_norms, marker='^', linestyle='-', color='g')\n", "axes[2].set_title(\"Gradient Norm\")\n", "axes[2].set_xlabel(\"Step\")\n", "axes[2].set_ylabel(\"Grad Norm\")\n", "axes[2].grid(True)\n", "\n", "plt.tight_layout()\n", "\n", "# 保存图片到本地\n", "image_path = \"training_dashboard.png\"\n", "plt.savefig(image_path)\n", "plt.show()\n", "\n", "# 触发浏览器下载原始 json 和图片\n", "print(\"\\n✅ 正在下载原始 trainer_state.json 和 综合训练指标图表...\")\n", "files.download(state_path)\n", "files.download(image_path)" ], "execution_count": null, "outputs": [] }, { "cell_type": "markdown", "metadata": { "id": "0f63c856" }, "source": [ "### 2. 下载模型权重到本地\n", "将模型文件夹打包为 ZIP 压缩包,并触发浏览器下载。" ] }, { "cell_type": "code", "metadata": { "id": "b686e5af" }, "source": [ "import json\n", "import shutil\n", "from google.colab import files\n", "\n", "with open(\"_train_config.json\") as f:\n", " cfg = json.load(f)\n", "\n", "suffix = \"_comparison\" if cfg[\"COMPARISON_TRAIN\"] else \"\"\n", "safetensors_dir = f\"safetensors_cocreator_qwen35_0.8b{suffix}\"\n", "zip_filename = f\"{safetensors_dir}.zip\"\n", "\n", "print(f\"正在将 {safetensors_dir} 打包为 {zip_filename},请稍候...\")\n", "shutil.make_archive(safetensors_dir, 'zip', safetensors_dir)\n", "print(\"✅ 打包完成!即将开始下载...\")\n", "\n", "files.download(zip_filename)\n" ], "execution_count": null, "outputs": [] }, { "cell_type": "markdown", "source": [ "### 3. 早停合并权重" ], "metadata": { "id": "NVfi0Rz1DkQk" } }, { "cell_type": "code", "source": [ "merge_content = \"\"\"\n", "from unsloth import FastVisionModel\n", "\n", "model, tokenizer = FastVisionModel.from_pretrained(\n", " model_name=\"outputs_qwen35_0.8b_cocreator/checkpoint-105\",\n", " max_seq_length=2048,\n", " load_in_4bit=False,\n", ")\n", "\n", "model.save_pretrained_merged(\n", " \"safetensors_cocreator_qwen35_0.8b\", # 和原脚本完全一致\n", " tokenizer,\n", " save_method=\"merged_16bit\"\n", ")\n", "\"\"\"\n", "\n", "with open(\"merge.py\", \"w\") as f:\n", " f.write(merge_content)\n", "\n", "!uv run python merge.py" ], "metadata": { "id": "4F6MyFpO-n3W" }, "execution_count": null, "outputs": [] } ] }