ARTICLE DETAIL

资讯详情

深耕网站建设、视觉设计与SEO优化的一线实战洞察。

使用 JAX/Flax 微调 ViT 进行图像分类:基于 Hugging Face Transformers 的完整实战指南

使用 JAX/Flax 微调 ViT 进行图像分类:基于 Hugging Face Transformers 的完整实战指南 推理引擎大模型【免费下载链接】FlexGenRunning large language models on a single GPU for throughput-oriented scenarios.项目地址https://gitcode.com/gh_mirrors/fl/FlexGen点击查看免费下载导读本指南基于 Hugging Face Transformers 仓库中的 JAX/Flax 视觉示例benchmark/third_party/transformers/examples/flax/vision/讲解如何利用 JAX/Flax 后端在 imagenette 数据集上微调 Vision TransformerViT完成图像分类任务。JAX/Flax 允许将纯函数跟踪并编译为高效、融合的加速器代码可在 GPU 与 TPU 上运行通过本文你将掌握数据集准备、训练脚本参数调优、训练/评估循环原理以及将模型推送至 Model Hub 的完整流程。JAX/Flax 图像分类示例概览本示例以 run_image_classification.py 为核心脚本使用FlaxAutoModelForImageClassification自动加载与微调 ViT 模型。从 modeling_flax_auto.py 的自动映射表可以看到JAX/Flax 后端当前支持beitFlaxBeitForImageClassification与vitFlaxViTForImageClassification两类图像分类模型意味着本脚本不仅限于 ViT也能用于 BeIT 的微调。JAX/Flax 的核心理念与 PyTorch 训练模式差异显著纯函数与不可变模型JAX 模型是不可变的参数更新以纯函数方式完成params - loss - grad - new_params这为模型并行提供了简单高效的实现基础即时编译通过 JAX 的jit将纯函数跟踪并编译为融合的加速器代码设备并行pmap在多个设备GPU/TPU core上并行执行计算脚本检测到多少设备就自动使用多少设备分布式训练开箱即用。注当前 JAX/Flax 后端没有类似Trainer的高级抽象所有示例都包含显式的训练循环这也是理解本脚本源码价值的关键点。准备数据集示例使用 imagenette 数据集进行微调。Imagenette 是 Imagenet 的一个子集只包含 10 个易于分类的类别tench丁鲷鱼、English springer英国史宾格犬、cassette player卡带播放机、chain saw链锯、church教堂、French horn法国号、garbage truck垃圾车、gas pump加油泵、golf ball高尔夫球、parachute降落伞。下载并解压数据wget https://s3.amazonaws.com/fast-ai-imageclas/imagenette2.tgz tar -xvzf imagenette2.tgz解压后会生成imagenette2目录包含train和val两个子目录每个类别在子目录下再分目录存放图片。训练脚本期望的目录结构为按类别分目录的标准 ImageFolder 格式root/dog/xxx.png root/dog/xxy.png root/dog/[...]/xxz.png root/cat/123.png root/cat/nsdf3.png root/cat/[...]/asd932_.png其中root即脚本参数中的--train_dir或--validation_dir所指的目录类别名由各子目录名自动推断脚本通过torchvision.datasets.ImageFolder读取并利用train_dataset.classes决定分类头的num_labels。安装依赖训练脚本依赖 requirements.txt 中声明的环境jax0.2.8 jaxlib0.1.59 flax0.3.5 optax0.0.8 torch1.9.0cpu torchvision0.10.0cpu注意脚本虽以 JAX/Flax 为核心但数据加载与预处理复用 torchvisionImageFolder、transforms与DataLoader因此 PyTorch 的 CPU 版本即可满足数据侧需求。若在 GPU 上运行还需参照官方 JAX 安装指南根据 CUDA 与 CuDNN 版本安装对应的jaxlib。训练模型运行以下命令即可微调模型python run_image_classification.py \ --output_dir ./vit-base-patch16-imagenette \ --model_name_or_path google/vit-base-patch16-224-in21k \ --train_dirimagenette2/train \ --validation_dirimagenette2/val \ --num_train_epochs 5 \ --learning_rate 1e-3 \ --per_device_train_batch_size 128 --per_device_eval_batch_size 128 \ --overwrite_output_dir \ --preprocessing_num_workers 32 \ --push_to_hub原文档记录在单 GPU 上该命令约7 分钟即可完成训练验证准确率可达99%。这一表现得益于预训练权重google/vit-base-patch16-224-in21k已在 ImageNet-21k 上完成预训练下游 10 分类任务只需少量 epoch 即可收敛高学习率1e-3配合线性预热 线性衰减调度适合迁移微调场景大批量 128 在 JAX 编译与设备并行下开销可控。训练参数详解脚本使用HfArgumentParser解析三组参数ModelArguments模型、DataTrainingArguments数据、TrainingArguments训练。除命令行传参外也支持将全部参数写入 JSON 文件后以python run_image_classification.py args.json方式传入源码 run_image_classification.py 中对单参数且以.json结尾的场景调用parse_json_file。模型参数ModelArguments参数默认值说明--model_name_or_pathNone预训练 checkpoint用于权重初始化不设置则从零训练--model_typeNone从零训练时指定模型类型vit、beit等--config_nameNone与模型名不同的预训练 config 名称或路径--cache_dirNone预训练模型下载缓存目录--dtypefloat32权重初始化与训练的浮点格式可选[float32, float16, bfloat16]--use_auth_tokenFalse使用huggingface-cli login生成的 token 访问私有模型数据参数DataTrainingArguments参数默认值说明--train_dir必填训练数据根目录每个类别一个子目录--validation_dir必填验证数据根目录结构同训练集--image_size224输入图片分辨率ViT patch 尺寸为 16 时需能被 224 整除--max_train_samplesNone调试用截断训练样本数--max_eval_samplesNone调试用截断评估样本数--overwrite_cacheFalse覆盖缓存的数据集--preprocessing_num_workersNone数据预处理进程数示例命令中设为 32 以加速 CPU 侧预处理训练参数TrainingArguments参数默认值说明--output_dir必填预测结果与 checkpoint 输出目录--overwrite_output_dirFalse覆盖输出目录内容若output_dir指向 checkpoint 目录可用于继续训练--do_train/--do_evalFalse是否执行训练 / 评估--per_device_train_batch_size8每个 GPU/TPU core/CPU 的训练批大小--per_device_eval_batch_size8每个设备的评估批大小--learning_rate5e-5AdamW 初始学习率--weight_decay0.0AdamW 权重衰减--adam_beta1/--adam_beta20.9/0.999AdamW 的 Beta 系数--adam_epsilon1e-8AdamW 的 Epsilon--adafactorFalse是否用 Adafactor 替代 AdamW--num_train_epochs3.0总训练轮数--warmup_steps0线性预热步数--logging_steps500每 X 更新步打印日志--save_steps500每 X 更新步保存 checkpoint--eval_stepsNone每 X 步执行一次评估--seed42随机种子--push_to_hubFalse训练后是否上传模型至 Hub--hub_model_idNone与本地output_dir同步的 Hub 仓库名--hub_tokenNone推送至 Hub 的认证 token提示脚本默认per_device_train_batch_size为 8示例命令显式调至 128若显存不足可酌情降低。实际全局批大小 每设备批大小 ×jax.device_count()见源码 run_image_classification.py。源码级解析数据增强与训练循环数据增强策略脚本通过 torchvision 完成预处理源码 run_image_classification.py训练与评估采用不同的增强策略训练集RandomResizedCrop(image_size)RandomHorizontalFlip()ToTensor() 归一化Normalize(mean[0.5,0.5,0.5], std[0.5,0.5,0.5])验证集Resize(image_size)CenterCrop(image_size)ToTensor() 相同归一化。源码注释特别指出这里使用默认预处理若要追求极致准确率应针对任务精心调优变换组合。归一化均值/标准差均取 0.5即像素值从[0,1]映射到[-1,1]这与 ViT 预训练时的输入分布一致。学习率调度与优化器create_learning_rate_fn源码 run_image_classification.py实现线性预热 线性衰减双段调度预热段optax.linear_schedule(init_value0.0, end_valuelearning_rate, transition_stepsnum_warmup_steps)衰减段从learning_rate线性降至 0跨越num_train_steps - num_warmup_steps步两段通过optax.join_schedules拼接边界为num_warmup_steps。优化器为optax.adamw参数b1、b2、eps、weight_decay均来自命令行源码 run_image_classification.py。训练步与设备并行核心训练逻辑源码 run_image_classification.pytrain_step内部用jax.value_and_grad(compute_loss)同时计算损失与梯度梯度经jax.lax.pmean(grad, batch)跨设备求平均实现数据并行下的梯度同步p_train_step jax.pmap(train_step, batch, donate_argnums(0,))将训练步编译为多设备并行版本donate_argnums允许原地更新状态以节省内存TrainState继承flax.training.train_state.TrainState并额外携带dropout_rng通过jax_utils.replicate复制到各设备dropout_rng用shard_prng_key分片保证每设备独立的随机流评估步同样pmap并行化并用pad_shard_unpad处理最后一个 batch 不足整设备批大小的填充问题。每个 epoch 结束后脚本在主进程jax.process_index() 0调用model.save_pretrained(output_dir, paramsparams)保存 checkpoint若开启--push_to_hub则同步推送至 Hub源码 run_image_classification.py。将模型上传至 Model Hub所有 JAX/Flax 示例脚本都支持通过--push_to_hub参数自动上传最终模型至 Hugging Face Model Hub。具体行为未指定--hub_model_id时仓库名自动生成你的用户名 / output_dir 的文件夹名例如用户名sgugger、输出目录~/tmp/test-mrpc则生成sgugger/test-mrpc指定--hub_model_id时需填写完整仓库名含用户名例如--hub_model_id sgugger/finetuned-bert-mrpc如需上传到组织用组织名替换用户名即可使用前提本地需先执行huggingface-cli login登录 Hugging Face 账号或通过--hub_token传入认证 token注意--output_dir必须是全新目录或远端仓库的本地克隆脚本在push_to_hub模式下会调用Repository(output_dir, clone_fromrepo_name)同步该目录。运行环境与扩展阅读本示例设计为在 Cloud TPU 上高效运行同时也支持单卡与多卡 GPUJAX 的 GPU 安装取决于 CUDA/CuDNN 版本需按官方指南选择对应jaxlib若想快速验证脚本流程可参考 test_flax_examples.py 中的测试模式如--max_train_samples、--max_eval_samples截断数据、降低 epoch 与 batch以缩短迭代周期JAX/Flax 相关的更多示例语言建模、文本分类、问答等位于 benchmark/third_party/transformers/examples/flax/其 README.md 提供了 JAX/Flax 生态、TPU/GPU 部署与 Hub 上传的通用说明脚本入口、参数解析与训练循环的完整实现请参阅 run_image_classification.py。结语通过本文你可以完整复现「下载 imagenette → 组织 ImageFolder 目录 → 一键微调 ViT → 推送 Hub」的 JAX/Flax 图像分类流程。深入源码后可以看到torchvision 负责数据侧、JAX/Flax 负责计算侧的分工模式以及pmap设备并行、线性预热/衰减调度、不可变 TrainState 更新等关键机制。这一示例也是理解 Transformers 库 JAX/Flax 后端训练范式的最佳起点——掌握它之后阅读其他 Flax 示例如语言建模、文本分类将事半功倍。赞分享推理引擎大模型【免费下载链接】FlexGenRunning large language models on a single GPU for throughput-oriented scenarios.项目地址https://gitcode.com/gh_mirrors/fl/FlexGen点击查看免费下载相关推荐使用 Transformers 微调 ViT 图像分类模型从 Food-101 数据集到 Hugging Face Hub 的完整实战指南使用 Transformers 微调 ViT 图像分类模型从 Food 101 数据集到 Hugging Face Hub 的完整实战指南 本篇技术指南基于人工智能AI 技能/插件大模型AI 评测CANN/asc-devkitAscend C SIMD API存储非对齐数据接口asc_storeunalign_post_postupdate 产品支持情况 | 产品 | 是否支持 | | : | : :| | Ascend 950PR/人工智能深度学习算子库CANNAscend使用 timm 模型与 Hugging Face Trainer 进行图像分类微调TimmWrapper 实战指南使用 timm 模型与 Hugging Face Trainer 进行图像分类微调TimmWrapper 实战指南 本指南聚焦于 huggingface vi人工智能AI 技能/插件大模型AI 评测上一篇如何彻底解决GTA圣安地列斯的技术问题SilentPatch修复方案深度解析下一篇Webiny-js本地化功能全解析多语言内容管理与国际化部署创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表