尧图建网站 尧图建网站 YAOTU WEB BUILD 免费咨询
ARTICLE DETAIL

资讯详情

深耕网站建设与建站编程的一线实战洞察。

Qwen3.5-0.8B进行微调实现NBA球员图片识别(1)

Qwen3.5-0.8B进行微调实现NBA球员图片识别(1) # 教 Qwen3.5-0.8B 认识 30 位 NBA 球员 --- ## 1. 环境 | 组件 | 版本/说明 | |---|---| | OS | WindowsPowerShell | | Python | 3.11.9venv-qwen35 虚拟环境 | | PyTorch | 2.5.1cu121CUDA 版非 cpu 版 | | 硬件 | 4GB 显存 GPU 16GB RAM | | 框架 | Transformers 5.xdev 版 LLaMA-Factory PEFT | | 模型 | Qwen/Qwen3.5-0.8B原生多模态 | powershell py -3.11 -m venv venv-qwen35 .\venv-qwen35\Scripts\activate pip install torch torchvision --index-url https://download.pytorch.org/whl/cu1212. 结果指标数值验证集识别准确率127/150 84.7%训练集510 张见过train_loss 0.1952 → 最低 0.0002验证集150 张从没见过eval_loss 0.0432最优 checkpoint-300满分球员5/5 全对12 位库里、詹姆斯、杜兰特、字母哥、约基奇、东契奇、恩比德、伦纳德、Cade、香蕉船…最弱球员Booker 2/5、Wembanyama 2/5结论4GB 显存 LoRA只训 练0.8% 参数 660 张自建图片 →模型基本学会了认NBA球员。3. 问题定义输入单张球员照片任意尺寸自动缩放输出球员英文名如LeBron James形式化图片→ViT视觉token→LLM名字token图片 \xrightarrow{ViT} 视觉token \xrightarrow{LLM} 名字token图片ViT​视觉tokenLLM​名字token指标验证集逐张生成 → 与标注比对 → 全名精确匹配准确率4. 逐步实现Step 1 环境搭建# experiments/00_check_gpu.pyimporttorchprint(torch.__version__)# 2.5.1cu121 ← 必须带 cu121 而不是 cpuprint(torch.cuda.is_available())# Trueprint(torch.cuda.get_device_name(0))Step 2 架构解剖结构原理首次加载fromtransformersimportQwen3_5ForConditionalGeneration,AutoProcessor modelQwen3_5ForConditionalGeneration.from_pretrained(Qwen/Qwen3.5-0.8B,dtypetorch.bfloat16,device_mapauto)⚠️坑 4AutoModelForCausalLM只加载语言部分没有 visualgenerate 时拒收pixel_values✅ 必须用专用类Qwen3_5ForConditionalGenerationconfig.architectures 里确认顶层结构实测[model] 752.4 M ← 语言模型主体 [lm_head] 254.3 M ← 输出层与词嵌入共享权重24 层配方官方 实测确认6 × [ 3×(Gated DeltaNet→FFN) 1×(Gated Attention→FFN) ]️ 图片ViT 视觉编码器视觉 token 序列分辨率自适应 146~2453 个 文本Token Embedding248320 × 1024|vision_start|…|image_pad|…|vision_end|Early Fusion 统一序列18× Gated DeltaNetconv1d 接受门a 遗忘门b 输出门z6× Gated AttentionRoPE 标准注意力LM Head 名字输出Gated DeltaNet 本质对应LSTM 记忆点in_proj_a ≈ LSTM 输入门接受新信息 in_proj_b ≈ LSTM 遗忘门压缩旧状态 in_proj_z ≈ LSTM 输出门 conv1d 局部 n-gram 增强对比标准 attention不维护 attention matrix改维护压缩隐藏状态递推式长上下文262K tokens 原生显存不爆炸。Step 3 数据采集与构建采集策略30 球员 × 22 张 660 张双源互补来源每人张数特点维基百科词条主图1-2高清官方照人脸清晰锚点DuckDuckGo 图片搜索20多样性不同球衣、场景、角度# scrape_nba.py 核心fromduckduckgo_searchimportDDGS resultsDDGS().images(keywordsf{player}basketball,max_results30)foriteminresults:download(item[image],fdata/nba/images/{slug}/{i:02d}.jpg)数据格式LLaMA-Factory sharegpt 多模态标准{messages:[{role:user,content:image\n这张照片里的是哪位NBA球员请只回答球员名字。},{role:assistant,content:LeBron James}],images:[data/nba/images/lebron_james/01.jpg]}切分每位球员 17 训练 5 验证 → 510 train 150 val验证集严格隔离模型从没见过Step 4 LoRA 微调LLaMA-Factory### modelmodel_name_or_path:Qwen/Qwen3.5-0.8Bimage_max_pixels:262144### methodfinetuning_type:loralora_rank:16lora_alpha:32lora_target:allfreeze_vision_tower:truefreeze_multi_modal_projector:true### trainbf16:trueper_device_train_batch_size:1gradient_accumulation_steps:8learning_rate:2.0e-4num_train_epochs:5gradient_checkpointing:truedataloader_num_workers:0### datasetdataset:nba_traineval_dataset:nba_valtemplate:qwen3_vl_nothinkllamafactory-cli train train_nba.yaml得到Number of trainable parameters 10,822,656 ← 只训 0.8% 参数 {loss: 3.842 → 0.5619 → 0.3178 → 0.1777 → ... → 0.0002} eval_loss: 0.2068 →... →0.0432 ← 一路下降 Total optimization steps 320 ← 5 epochs × 64 Training completed in 2h07mStep 5 评估逐张测试# eval_nba.py 核心modelPeftModel.from_pretrained(base_model,checkpoints/lora-nba)oneprocessor.apply_chat_template(# return_dictTrue 是命门messages,tokenizeTrue,add_generation_promptTrue,return_dictTrue,return_tensorspt)inputs{k:v.to(model.device)fork,vinone.items()}outmodel.generate(**inputs,max_new_tokens32,do_sampleFalse)# 解码兼容 2D/3Dout[:, input_len:] 或 out[:, 0, input_len:]实测结果150 张验证图 →84.7%错误集中在 Booker/Wembanyama 这类特征相似或训练图偏少的球员5. 对比思考选择理由优化思路Qwen3.5-0.8B4G 显存能跑 原生多模态 LoRA 后仍轻换 2B/4B需更多显存或量化冻结 ViT球员特征主要由 LLM 对齐学解冻后半段 ViT 可提精度显存↑每球员 22 张分类任务泛化Booker/Wemby 各补 10 张再练一轮5 epochs小数据分类任务需要多轮记住映射看 eval_loss 是否回升过拟合信号6. 总结多模态模型 ViT 编码器 桥 LLM——图片切片成 token 和文字 token 混进同一序列Early Fusion这也是为什么显存消耗随图片分辨率剧烈变化Gated DeltaNet 把 LSTM 门控思想写进并行注意力18 层压缩历史 6 层精修长上下文不爆显存LoRA 微调在 4G 显存完全可行0.8B 模型 bf161.6GB 冻结 ViT batch 1 训练成本仅 0.8% 参数数据集非常重要路径基准CWD、格式sharegpt JSONL、切分验证隔离——错误率最高的一环7. 参考Qwen3.5 官方仓库github.com/QwenLM/Qwen3.8含 Qwen3.5 系列信息Qwen3.5-0.8B 模型卡huggingface.co/Qwen/Qwen3.5-0.8BLLaMA-Factorygithub.com/hiyouga/LLaMA-Factoryexamples/train_lora/qwen3vl_lora_sft.yaml 为配置蓝本duckduckgo-searchgithub.com/deedy5/duckduckgo_searchPEFTgithub.com/huggingface/peft
返回列表