Skip to content

Repository files navigation

logo_detailed_prompt - 详细提示词到 SVG 徽标的训练项目

本仓库包含用于 Part B 作业的数据、奖励函数、LoRA 微调脚本、自评脚本、训练结果和报告。

数据集用于训练一个小模型:输入为详细的视觉提示词,输出为完整的 SVG 徽标。目标 SVG 由 Claude Sonnet 根据提示词生成。

数据文件

文件 行数 内容
train.jsonl 219 训练样本
valid.jsonl 17 验证样本

每一行都是 chat 格式样本:

{"messages": [
  {"role": "system",    "content": "<SVG 设计器指令>"},
  {"role": "user",      "content": "<详细视觉提示词>"},
  {"role": "assistant", "content": "<完整 <svg>...</svg> 文档>"}
]}
  • 输入是 user 字段里的详细提示词。
  • 目标是 assistant 字段里的完整 SVG 文档,使用 viewBox="0 0 256 256"
  • 训练时只对 assistant 的 SVG token 计算 loss,system/user 部分需要 mask。

数据来源

  • 原始数据包含 275 条生成记录。删除不完整记录并修复不合法 SVG 后,剩余 253 条有效的“详细提示词到 Sonnet-SVG”配对。
  • 其中 17 条作为私有测试集保留,不包含在本仓库中;公开部分共 236 条。
  • 本仓库不包含 raw-query augmentation 行,只包含详细提示词配对。

项目文件

路径 作用
reward.py 作业提交用奖励函数,返回总分和分项诊断。
student_kit/reward.py 对顶层 reward.py 的兼容包装,便于脚本导入。
train_peft.py 基于 PEFT 的 LoRA 训练脚本,显式 mask prompt token。
train_config.yaml 本次训练使用的超参数配置。
student_kit/eval_self.py 基座模型与 LoRA adapter 的生成式自评脚本。
results.dry_run.json 用验证集金标准 SVG 做的 reward 校准结果,不是最终模型结果。
results.json 基座模型与 LoRA adapter 在验证集上的自评结果。
report.md 中文实验报告,包含 reward 设计、训练设置、结果和分析。
adapter/ 已训练的 LoRA adapter。

环境准备

建议在有 CUDA 的 Python 环境中安装依赖:

pip install -r requirements.txt
pip install modelscope

从 ModelScope 下载 Gemma 3 270M instruction-tuned 基座模型:

modelscope download --model google/gemma-3-270m-it --local_dir ./gemma3-270m

本次本地实验使用的是 RTX 4060 Laptop GPU。由于 8GB 显存限制,训练配置中关闭了训练中的验证 loss,改为训练后运行生成式自评。

训练

运行:

python train_peft.py --config train_config.yaml

训练完成后会在 adapter/ 中保存 LoRA adapter。作业要求的核心文件是:

  • adapter/adapter_config.json
  • adapter/adapter_model.safetensors

自评

运行固定贪心解码的基座模型 vs LoRA adapter 自评:

python student_kit/eval_self.py \
  --base-model ./gemma3-270m \
  --adapter ./adapter \
  --valid valid.jsonl \
  --output results.json

不加载模型、只验证 reward 和数据管线时,可以运行:

python student_kit/eval_self.py --dry-run-targets --valid valid.jsonl --output results.dry_run.json

本次结果

本次 LoRA 训练完成,但验证集 reward 没有超过基座模型:

模型 平均 reward
基座 Gemma 3 270M 0.157955
LoRA 微调模型 0.100000
Delta -0.057955

主要原因是 LoRA 模型学会了直接输出 <svg ...> 开头和部分颜色/元素风格,但大多数输出没有闭合 </svg>,因此无法被 reward 解析为完整 SVG。详细分析见 report.md

最终提交清单

  • adapter/adapter_config.json
  • adapter/adapter_model.safetensors
  • reward.py
  • train_config.yaml
  • results.json
  • report.md

About

No description, website, or topics provided.

Resources

Stars

0 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages