免费获取学习方案
ARTICLE DETAIL

资讯详情

深耕编程基础知识与建站技术分享的一线实战洞察。

基于 Flower 与 Whisper-tiny 的设备端联邦微调实战:Speech Commands 关键词识别(KWS)

基于 Flower 与 Whisper-tiny 的设备端联邦微调实战:Speech Commands 关键词识别(KWS) 基于 Flower 与 Whisper-tiny 的设备端联邦微调实战Speech Commands 关键词识别KWS【免费下载链接】flowerFlower: A Friendly Federated AI Framework项目地址: https://gitcode.com/GitHub_Trending/flo/flower本示例演示如何从预训练的 Whisper 模型出发通过 Flower 联邦学习框架构建一个 100 客户端规模的端到端设备端on-device下游微调流水线冻结 Whisper-tiny 的编码器参数仅学习一个不足 800K 参数的轻量分类头在 Google Speech Commands 数据集上完成关键词识别Keyword SpottingKWS任务。读完本文你将掌握集中式训练与联邦微调两种方案的完整跑通方法理解如何用GroupedNaturalIdPartitioner按说话人 ID 构造真实世界的非均衡联邦数据划分并能在 Simulation Engine 与 Deployment Engine 两种模式下运行同一套代码甚至将其部署到 Raspberry Pi 上。示例概览冻结编码器 轻量分类头本示例位于仓库的 examples/whisper-federated-finetuning 目录核心思路非常简洁取 openai/whisper-tiny 模型的编码器encoder部分冻结其全部参数作为语音特征提取器在其输出之上挂接一个轻量分类头用于把一段 1 秒语音分类到 12 个关键词类别之一整个联邦训练过程只聚合和更新这个分类头通信开销极小。从源码 model.py 可以看到分类头的具体结构Conv1d(1500, 128, kernel_size1)→ReLU→Flatten→Linear(128 * 384, num_classes)参数总量约 78 万README 中记为800K parameters集中式脚本打印的实际值为781964。由于分类头如此轻量客户端与服务端之间的通信成本被降到了最低非常适合网络受限的设备端场景。下游数据集采用 Google Speech Commandsv0.02用于关键词识别任务。本示例可以在三种模式下运行模式说明集中式训练Centralized传统 ML 训练方式所有数据对执行微调的节点可见联邦学习-模拟Simulation客户端是瞬时的 Python 进程被分配系统资源的一部分联邦学习-设备端On-device客户端是相互独立的实体各自可以运行在不同的设备上项目搭建与依赖安装克隆项目在 examples/whisper-federated-finetuning 目录之外可通过以下命令把本示例单独取出来git clone --depth1 https://gitcode.com/GitHub_Trending/flo/flower.git _tmp \ mv _tmp/examples/whisper-federated-finetuning . \ rm -rf _tmp \ cd whisper-federated-finetuning执行后将得到一个名为whisper-federated-finetuning的目录包含以下文件whisper-federated-finetuning ├── whisper_example │ ├── __init__.py │ ├── client_app.py # 定义你的 ClientApp │ ├── server_app.py # 定义你的 ServerApp │ ├── model.py # 定义模型与训练函数 │ └── dataset.py # 定义数据集及其处理流程 ├── centralized.py # 本示例的集中式版本 ├── preprocess.py # 预处理器用于把全部分区保存到磁盘 ├── pyproject.toml # 项目元数据、依赖与运行配置 └── README.md安装依赖与项目包在新建的 Python 环境中安装 pyproject.toml 定义的依赖以及whisper_example包本身pip install -e .pyproject.toml的核心依赖包括flwr[simulation]1.28.0、flwr-datasets[audio]0.5.0、transformers4.53.0与torch并声明了两个关键组件入口[tool.flwr.app.components] serverapp whisper_example.server_app:app clientapp whisper_example.client_app:app集中式训练先建立性能基线本节描述不借助联邦学习直接微调Whisper-tiny的方法即每一轮训练都使用完整的训练集。想直接体验 Flower 联邦版本可以跳到下一节。运行集中式训练python centralized.py --compile # 如果你使用 pytorch 2.0请不要加 --compile 参数 # 脚本会在每个 epoch 结束后保存分类头的 checkpoint # checkpoint 遵循如下命名风格: classifier_val_accuracy.pt # 通过如下方式加载某个 checkpoint 继续训练或推理: python centralized.py --checkpoint my_checkpoint.pt首次运行代码时SpeechCommands数据集会被下载并预处理通过 API大约需要 40 分钟缓存在~/.cache/huggingface/datasets/speechcommands占用约 83GB 磁盘空间。后续运行不再需要重复这一预处理过程。从 centralized.py 源码可以看到集中式训练的实现细节支持--epochs默认 3 轮、--compilePyTorch 2.0 模型编译加速与--checkpoint三个命令行参数使用get_encoding_fn对 train/validation/test 三个划分做同样的特征编码调用prepare_silences_dataset以 10% 比例生成silence类别的训练样本使用construct_balanced_sampler构造类别均衡采样器确保每个 batch 内各类别样本数量大致相当优化器为torch.optim.SGD(classifier.parameters(), lr0.001)只训练分类头编码器保持encoder.eval()冻结状态。README 给出的预期结果是2 个 epoch 内验证准确率可超过 95%在 RTX 3090Ti 上每个 epoch 约需 3 分 30 秒最终测试集稳定达到 97% 以上。参考日志如下... classifier_head_params 781964 Initial (loss, acc): loss 0.04124763025785586, accuracy 0.03215788419154478 Epoch: 0 100%|████████████████████████| 84928/84928 [03:0500:00, 456.93it/s, avg_loss0.7269, avg_acc0.8282] VALIDATION --- loss 0.0051703976778501234, accuracy 0.9319775596072931 Epoch: 1 100%|████████████████████████| 84928/84928 [03:0700:00, 454.06it/s, avg_loss0.1588, avg_acc0.9629] VALIDATION --- loss 0.003613288299632327, accuracy 0.943097575636145 Epoch: 2 100%|████████████████████████| 84928/84928 [03:0600:00, 454.16it/s, avg_loss0.1208, avg_acc0.9675] VALIDATION --- loss 0.0022978041400064466, accuracy 0.9610298537367261 Training done... Evaluating test set. Loading best model TEST --- loss 0.001703281509680464, accuracy 0.9740286298568507联邦化数据准备按 speaker_id 划分非均衡分区集中式训练虽好但在许多真实场景下无法落地——因为训练数据必须保持在客户端本地即数据不可出域无法聚合到单一节点如服务器。这正是 Flower 联邦微调流水线的用武之地客户端在本地数据上训练分类头把更新上传到中心服务器服务器聚合后再分发给客户端进行下一轮联邦学习如此往复直至收敛。Speech Commands 的train分区是我们联邦学习所用的划分。先看它的统计信息from datasets import load_dataset sc_train load_dataset(speech_commands, v0.02, splittrain, tokenFalse) print(sc_train) # Dataset({ # features: [file, audio, label, is_unknown, speaker_id, utterance_id], # num_rows: 84848 # }) # 训练集由约 8.5 万条 1 秒音频片段构成来自 2112 位说话人 ids set(sc_train[speaker_id]) print(len(ids)) # 2113 # --- 1 是因为包含了一个 None 说话人用于构造 _silence_ 训练样本本示例使用 Flower Datasets 提供的GroupedNaturalIdPartitioner基于speaker_id对 SpeechCommands 数据集进行分区。具体做法是创建每组 5 个说话人的分组最终得到 422 个分组每组代表联邦中的一个节点/客户端每个speaker_id只出现在一个分组中。可以把每个分组想象成一个联邦学习节点其中包含多位用户/说话人——例如一个办公室里有多位员工在使用同一个关键词识别系统。从 dataset.py 源码可以看到数据加载的核心逻辑FederatedDataset以全局缓存方式只初始化一次GroupedNaturalIdPartitioner(partition_byspeaker_id, group_size5)负责分区随后对每个分区执行特征编码并按比例注入silence样本每个客户端的沉默样本数量为其样本数的 10% 按全局占比折算。由于并非所有speaker_id贡献了同样多的音频片段Speech Commands 数据集创建时的客观情况最终得到的数据分区大小并不均衡——这恰恰是真实世界中常见的现象。如果把每个客户端/节点的数据量画成柱状图结果如下你可以运行 visualize_labels.ipynb 来生成或调整这张图它使用了 Flower Datasets 的可视化工具。联邦微调流水线ServerApp 与 ClientApp 的分工Flower 为这套流水线构建的整体流程如下图所示每一轮包含四个步骤下发每轮开始时ServerApp把分类头的权重分发给一部分节点本地训练每个节点的ClientApp使用冻结的预训练 Whisper 编码器在自身数据上训练分类头回传本地训练完成后每个节点把更新后的分类头上传给ServerApp聚合Flower 的ServerApp通过 FedAvg 聚合各分类头得到新的全局分类头并在下一轮分享给节点。当然你也可以选择其他策略或实现自定义策略。ClientApp本地训练逻辑client_app.py 中的train函数完整展示了客户端行为从context.node_config读取partition-id分区编号从context.run_config读取num-classes、batch-size、disable-tqdm、compile-model等运行配置通过get_model(device, num_classes, compile_model)构建编码器与分类头并用收到的权重初始化分类头classifier.load_state_dict(msg.content[arrays].to_torch_state_dict())加载分区数据后如果样本数大于batch_size用construct_balanced_sampler构造类别均衡采样器数据太少时放弃采样器直接顺序取 batch采用 Adam 优化器lr0.001与CrossEntropyLoss调用train_one_epoch训练一个 epoch有一个值得注意的细节如果某个客户端的样本数不足以构成一个完整 batchlen(train_loader) 1则跳过训练并通过ConfigRecord({trained: run_training})标记该客户端未训练最终以ArrayRecord模型权重、MetricRecord训练指标和ConfigRecord是否训练标记组装回复消息。ServerApp聚合与中心评估server_app.py 定义了服务端逻辑从context.run_config读取num-server-rounds、num-classes、fraction-train初始化全局分类头参数同样只联邦分类头_, classifier get_model(cpu, num_classes, False)若central-eval开启服务端加载 validation/test 划分构造eval_fn在每轮聚合后评估全局模型最后一轮自动切换为测试集评估对应日志中的test_accuracy使用自定义的ExclusiveFedAvg继承自 FedAvg做聚合——它遍历所有客户端回复剔除那些因数据不足而未参与训练trainedFalse的客户端只聚合真正训练过的模型避免把未训练的权重混入全局平均。源码中还会打印{n}/{total} models included for aggregation.供观察每轮的参与情况训练结束后把最终分类头保存为final_model.pt。num_classes 12的由来Speech Commands 中所有未知关键词被统一映射为标签 11silence片段被映射为标签 10从而形成 0–11 共 12 个类别见 dataset.py 中的get_encoding_fn。使用 Simulation Engine 运行联邦实验同一套代码可以在模拟与部署两种模式下运行而无需改动代码。如果你是 Flower 新手建议先使用模拟模式因为它需要手动启动的组件更少。默认情况下flwr run会使用 Simulation Engine。运行配置定义在 pyproject.toml 的[tool.flwr.app.config]块中[tool.flwr.app.config] num-server-rounds 3 fraction-train 0.05 # 每轮采样 5% 的客户端422 的 5% 即约 21 个 num-classes 12 batch-size 8 compile-model false disable-tqdm true central-eval false remove-cols file,audio,label,is_unknown,speaker_id,utterance_id本示例按 422 个虚拟SuperNode设计即 2112 位说话人按 5 人一组分组的结果。Simulation Runtime 默认只支持 10 个节点因此首先需要调整配置本指南假定你的默认SuperLink连接是可用于模拟的不确定时可参考 How-to run Flower locally 指南flwr federation simulation-config \ --num-supernodes422 \ --client-resources-num-cpus4 \ --init-args-log-to-driverfalse # 设为 true 可开启模拟引擎的全部日志默认情况下模拟只在 CPU 上运行。在 MacBook Pro M2 上运行 3 轮 Flower 联邦约需 10 分钟前提是数据集已下载。推荐在 GPU 上运行。关于 Flower Simulation 的工作原理及如何让其使用 GPU可查阅相关文档。然后直接运行# 使用默认设置运行每轮从 422 个客户端中采样 21 个 flwr run . --stream结束时你会看到联邦指标汇总即本轮采样客户端的平均训练准确率与损失形如INFO : [SUMMARY] INFO : Run finished 3 round(s) in 564.50s INFO : History (metrics, distributed, fit): INFO : {train_accuracy: [(1, 0.637721849625075), INFO : (2, 0.8666815319504736), INFO : (3, 0.8912498749526644)], INFO : train_loss: [(1, 4.049714171341712), INFO : (2, 1.8473016127565092), INFO : (3, 2.5116721350250693)]} INFO :在 GPU 上运行客户端需要先在 Flower 配置文件中定义一个新的 SuperLink 连接为虚拟客户端分配 GPU 资源参见 Flower Simulation 文档中 Defining ClientApp Resources 一节。覆盖运行配置pyproject.toml中定义的ClientApp/ServerApp设置均可通过--run-config覆盖。例如# 运行 10 轮每轮采样 20% 的客户端 flwr run . --run-config num-server-rounds10 fraction-fit0.2性能预期README 给出的参考值仅 5 轮联邦训练全局模型即可达到约 97% 的验证准确率使用默认超参数训练 10 轮可达到 96% 的测试准确率。在 RTX 3090Ti 上每轮约需 40–50 秒取决于当轮采样客户端拥有的数据量。开启中心评估并训练 10 轮命令为flwr run . --run-config central-evaltrue num-server-rounds10进阶挑战如果觉得当前联邦设置不够有挑战性可以减小GroupedNaturalIdPartitioner创建的组大小如每组 3 人、2 人这会增加联邦中客户端/节点的数量使数据划分更碎片化、更不均衡。使用 Deployment Engine 部署到真实设备以下步骤概述了把本示例从 Simulation Engine 切换到 Deployment Engine 所需的少量代码改动。Deployment Engine 的入门指南含启用安全 TLS 与节点认证请查阅相关文档。与模拟模式相比运行完全相同的联邦流水线无需改动ServerApp设计只需对ClientApp的数据集加载逻辑略作调整在模拟模式下我们希望动态地让一个 Python 进程扮演某个特定客户端加载其对应分区在部署模式下我们希望同一个客户端进程与单个SuperNode绑定始终使用运行该SuperNode的机器上本地存储的自己的数据集。因此第一步是生成 N 份数据分区并分派给不同的SuperNode用 preprocess.py 完成具体分三步1. 保存数据分区运行两次下面的命令每次指定不同的分区 id。每次运行都会在当前目录生成一个形如partition_id的目录python preprocess.py --partition-id5从源码看preprocess.py会从pyproject.toml读取remove-cols与num-supernodes配置不带--partition-id运行时它会用multiprocessing.Pool并行处理并把全部 422 个分区保存到磁盘。2. 调整client_fn把whisper_example/client_app.py中的partition-id键改名为更有意义的local-dataset并把load_data调用替换为load_data_from_disk这样ClientApp就会使用启动SuperNode时指定的本地数据集from whisper_example.dataset import load_data_from_disk app.train() def train(msg: Message, context: Context): # ... # partition_id context.node_config[partition-id] # 注释掉 local_data context.node_config[local-data] # 新的一行 # 其余保持不变 # 把原来的 load_data 相关代码替换为 partition load_data_from_disk(local_data)对应的load_data_from_disk实现位于 dataset.py本质是对 HuggingFaceload_from_disk的一层封装。3. 让 SuperNode 可访问数据把第 1 步生成的目录复制到将要运行SuperNode的机器上例如使用scp安全传输。完成以上三步后即可用 Deployment Engine 运行联邦 Whisper 微调。假设SuperNode所在机器的 Python 环境中已安装全部依赖flwr、transformers、torch见pyproject.toml把SuperNode连接到运行中的联邦即正在运行的SuperLinkflower-supernode --superlinkSUPERLINK-IP:9092 \ --node-configlocal-datapath/to/local/partition4. 运行你的 Whisper 应用首先确保 Flower 配置文件中定义了 SuperLink 连接。可以通过flwr config list定位配置文件打开后若没有现成连接则新建一个例如[superlink.remote] address 127.0.0.1:9093 # 你的 superlink 的 IP:9093此处假设 superlink 在本地 insecure true # 如需启用 SSL 请查阅相关文档当SuperNodes连接上SuperLink后通过flwr run启动运行这次把连接指向remoteflwr run . remote在 Raspberry Pi 上进行联邦微调在 Raspberry Pi 上启动 FlowerSuperNode步骤与在任何其他要接入联邦的机器上完全一致。首先确保你的 Raspberry Pi 已正确配置需要Raspberry Pi 4 或 5。按本示例现有代码运行树莓派上的内存占用不超过 1.5GB。注意与前面章节不同树莓派客户端更适合使用PyTorch 1.13.1或更早于 PyTorch 2.0 的版本。尚未配置 Pi 的话可参考 examples/embedded-devices 示例中的 Setup your Pi 一节完成设置。其次在开发机如笔记本上按前面 Run with the Deployment Engine 一节的步骤生成并复制一份数据分区到树莓派。最后假设你有一台机器如笔记本上运行着SuperLink且树莓派可以访问到它例如在同一局域网内即可像前面一样启动SuperNodeflower-supernode --superlinkSUPERLINK-IP:9092 \ --node-configlocal-datapath/to/local/partition关键源码速览与扩展方向模型编码器来自openai/whisper-tiny仅取 encoder 并冻结分类头约 78 万参数见 model.py训练采用标准 PyTorch 流程model.eval()classifier.train()编码器输出在torch.no_grad()下前向计算只对分类头反向传播见 model.py。数据FederatedDatasetGroupedNaturalIdPartitioner按说话人分组12 类标签映射与 silence 样本增强详见 dataset.py。聚合ExclusiveFedAvg过滤未训练客户端后再执行 FedAvg见 server_app.py。配置所有可调参数集中在 pyproject.toml 的[tool.flwr.app.config]中可通过--run-config随时覆盖无需改动代码。如果想进一步深入可以尝试减小group_size以制造更严苛的非均衡联邦场景替换FedAvg为其他内置策略或自定义策略开启central-eval观察每轮全局模型在验证集/测试集上的表现。这套以冻结大模型 联邦轻量头为范式的流水线同样适用于其他语音、视觉或多模态预训练模型的设备端联邦微调场景。【免费下载链接】flowerFlower: A Friendly Federated AI Framework项目地址: https://gitcode.com/GitHub_Trending/flo/flower创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表