ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

LoRA进阶:状态微调与并行控制实现大模型低显存高效训练

LoRA进阶:状态微调与并行控制实现大模型低显存高效训练 大家好我是专注于AI模型微调与部署的技术博主。在尝试使用LoRALow-Rank Adaptation技术对大语言模型进行个性化定制时你是否也遇到过这样的困境模型参数稍微大一点显存就瞬间告急训练过程频繁中断或者只能被迫使用极小的批次大小导致训练效率低下尤其是在多任务并行或需要同时微调多个适配器的场景下显存瓶颈更是成为了拦路虎。本文将深入探讨一种进阶的LoRA微调策略——从传统的“权重微调”转向更高效的“状态微调”并结合并行控制技术实现一套真正意义上的低显存消耗方案。无论你是刚接触LoRA的新手还是希望优化现有微调流程的开发者都能从本文获得一套从理论到实践的完整闭环指南。我们将从核心概念入手逐步拆解代码实现并提供可直接复用的配置与避坑指南。1. 背景与核心概念为什么需要更高效的微调在深入技术细节之前我们有必要厘清几个关键概念理解当前微调技术面临的挑战与演进方向。1.1 LoRA 微调的本质与局限LoRA低秩适应是一种参数高效微调PEFT技术。其核心思想并非直接更新原始大模型通常称为基础模型或基座模型的庞大权重矩阵而是冻结这些权重并引入一对可训练的、低秩的适配器矩阵通常记为 A 和 B。在模型的前向传播过程中原始权重 W 与低秩增量 ΔW BA 相加共同参与计算。公式化表示为h Wx ΔWx Wx BAx。这种方法极大地减少了需要训练的参数数量通常只有原模型的0.1%-1%从而显著降低了存储和计算开销。然而传统的LoRA实现权重微调在训练时仍然需要将基础模型的所有参数加载到显存中因为前向和反向传播的计算图依赖于完整的模型结构。对于拥有数十亿甚至上百亿参数的大模型仅加载模型本身就可能占满高端显卡的显存留给优化器状态、梯度、激活值和批次数据的空间就非常有限了。1.2 权重微调 vs. 状态微调这是本文要解决的核心矛盾也是实现低显存的关键。权重微调Weight Fine-Tuning这是我们最熟悉的模式。在训练循环的每一步优化器直接更新可训练参数即LoRA的 A 和 B 矩阵的权重值。优化器如Adam需要为每一个可训练参数维护两个状态动量一阶矩估计和方差二阶矩估计。这意味着即使可训练参数很少优化器状态也会带来额外的显存开销虽然相比全参数微调已小很多但在极端显存受限或并行多适配器场景下仍不可忽视。状态微调State Fine-Tuning这是一种更激进的思路。它不再直接更新权重参数而是更新优化器的状态。具体来说我们可以固定LoRA权重A, B的初始值甚至可以是零然后让优化器去学习如何调整其内部状态如Adam的动量和方差使得在推理时这些被“调整过的状态”能引导模型产生我们期望的输出。这听起来有些抽象但其优势在于优化器状态本身可能具有不同的、更高效的参数化方式或者可以在不同层、不同任务间共享从而潜在地实现更高的参数效率和更低的显存占用。一种简单的理解是它试图学习一个更好的“优化轨迹”或“更新规则”而非最终的权重点。1.3 并行控制下的低显存需求在实际应用中我们常常面临并行控制的需求多任务学习同时为同一个基础模型训练多个不同的LoRA适配器例如一个用于代码生成一个用于客服对话。超参数搜索并行运行多个具有不同超参数如学习率、秩大小的训练实验。集成学习训练多个LoRA适配器并进行集成。在“权重微调”模式下并行运行N个任务意味着需要在显存中同时保存N份基础模型参数和N份优化器状态显存消耗几乎是线性增长的。而“状态微调”结合一些并行控制策略如梯度检查点、模型并行、优化器状态卸载有望打破这种线性增长实现亚线性甚至常数的显存开销增长。2. 环境准备与版本说明为了完整复现后续的实战案例我们需要搭建以下环境。请注意版本号是示例核心是思路请根据你的实际环境进行调整。核心环境操作系统Ubuntu 20.04 LTS 或更高版本Windows WSL2 也可行但本文以Linux命令为例。Python3.8 或 3.9。推荐使用 conda 或 venv 创建独立的虚拟环境。PyTorch1.12 2.0 更佳。需与CUDA版本匹配。CUDA11.7 或 11.8根据你的GPU驱动选择。深度学习框架我们将主要使用 Hugging Face 的transformers和peft库。Python 包依赖创建一个requirements.txt文件内容如下torch2.0.0 transformers4.35.0 peft0.7.0 datasets2.14.0 accelerate0.25.0 # 用于分布式和混合精度训练 bitsandbytes0.41.0 # 可选用于4/8-bit量化进一步节省显存 trl0.7.0 # 可选用于RLHF等进阶训练 scipy sentencepiece protobuf安装命令# 创建并激活虚拟环境以conda为例 conda create -n lora_adv python3.9 conda activate lora_adv # 安装PyTorch请访问 https://pytorch.org/ 获取对应你CUDA版本的命令 # 例如 pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu118 # 安装其他依赖 pip install -r requirements.txt项目结构建议lora_advanced_tuning/ ├── configs/ # 存放训练配置 │ └── train_config.yaml ├── scripts/ # 存放训练、评估脚本 │ ├── train_weight.py # 传统权重微调脚本 │ └── train_state.py # 状态微调实验脚本 ├── models/ # 存放基座模型和微调后的LoRA权重 ├── data/ # 训练数据集 ├── outputs/ # 训练输出日志、检查点 └── utils/ # 工具函数 └── parallel_control.py3. 核心原理与方案拆解本节将深入讲解“状态微调”和“并行控制”是如何具体运作并节省显存的。3.1 低显存 LoRA 的传统优化手段在进入新方案前先回顾并整合已有的低显存技术它们是构建新方案的基础梯度检查点Gradient Checkpointing用计算时间换显存。只保留关键层的激活值非关键层的激活在反向传播时重新计算可以大幅减少激活值占用的显存。混合精度训练AMP使用torch.cuda.amp。让模型权重、激活和梯度使用float16半精度而优化器状态保持在float32单精度在几乎不影响精度的情况下减少显存占用和加速计算。4/8-bit 量化bitsandbytes使用bitsandbytes库加载模型将模型权重量化为4位或8位整数INT8/INT4进行存储和计算仅在前向传播时反量化为float16。这是目前节省模型参数显存最有效的方法之一。优化器状态卸载CPU Offloading将优化器状态、梯度甚至模型参数的一部分卸载到CPU内存仅在需要时传输到GPU。accelerate库提供了cpu_offload功能。3.2 状态微调的实现思路状态微调不是一个有标准API的现成功能而是一种设计模式。其一种实现路径如下定义可学习状态我们不直接学习权重矩阵A和B而是学习一组元参数Meta-Parametersθ。这些θ参数的数量远小于A和B。状态到权重的映射设计一个函数f将元参数θ和当前训练步骤t或其他上下文信息映射为当前步骤的权重增量ΔW_t。即ΔW_t f(θ, t)。这个f可以是一个简单的线性层一个小型神经网络甚至是一个查找表。优化目标损失函数L的计算依赖于由f(θ, t)生成的动态权重ΔW_t。我们计算损失关于元参数θ的梯度∇_θ L并用它来更新θ。推理阶段训练完成后我们保存的是元参数θ和函数f。在推理时对于给定的输入我们可以使用训练好的θ和f通常取最终或平均的权重映射来生成固定的ΔW然后与基础模型权重合并。节省显存的点元参数θ的维度极小因此其对应的优化器状态也极小。同时由于ΔW_t是动态生成的我们不需要在内存中一直保存所有A和B矩阵的多个副本对于并行任务只需保存一份θ和对应的生成函数即可。3.3 并行控制策略结合状态微调我们可以设计以下并行控制策略来管理多任务共享基础模型所有并行任务共享同一份基础模型参数在GPU显存中只存一份。这是LoRA的天然优势必须充分利用。任务特定的元参数每个并行任务i拥有自己独立的、小型的元参数集θ_i。序列化执行与显存复用利用accelerate或自定义上下文管理器在一个GPU上顺序执行多个任务的前向/反向传播。在每个任务计算完成后立即释放该任务独有的计算图、激活值和梯度只保留其微小的优化器状态在状态微调下这就是θ_i的优化器状态。这样峰值显存占用 ≈基础模型单个任务的计算开销N * 小尺寸优化器状态而不是N * (基础模型 计算开销)。优化器状态CPU卸载将θ_i的优化器状态也卸载到CPU仅在更新参数时同步到GPU可以进一步减少GPU显存压力。4. 完整实战案例基于 Qwen 模型的状态微调实验我们将以 Qwen-7B-Chat 模型为例展示一个简化的状态微调实现。请注意这是一个概念验证性的示例旨在阐明思路。4.1 项目结构与数据准备首先准备一个简单的指令微调数据集。我们使用datasets库加载一个示例数据集并格式化为对话形式。# scripts/data_prepare.py from datasets import load_dataset import json # 示例使用 Alpaca 格式的数据 def prepare_dataset(): # 这里可以替换成你自己的数据集 dataset load_dataset(yahma/alpaca-cleaned, splittrain[:100]) # 取100条做演示 def format_alpaca_to_conversation(example): # 将Alpaca (instruction, input, output) 格式化为 Qwen 的对话格式 conversation [ {role: system, content: You are a helpful assistant.}, {role: user, content: f{example[instruction]}\n{example[input]} if example[input] else example[instruction]}, {role: assistant, content: example[output]} ] # 转换为 transformers 训练器接受的字符串格式 # 实际使用时应使用 tokenizer.apply_chat_template return {text: json.dumps(conversation, ensure_asciiFalse)} formatted_dataset dataset.map(format_alpaca_to_conversation, remove_columnsdataset.column_names) formatted_dataset formatted_dataset.train_test_split(test_size0.1) return formatted_dataset if __name__ __main__: ds prepare_dataset() ds[train].to_json(./data/train.jsonl, orientrecords, linesTrue) ds[test].to_json(./data/test.jsonl, orientrecords, linesTrue) print(数据集准备完成。)4.2 实现状态微调适配器层我们创建一个自定义的 PEFT 层实现状态微调的逻辑。# utils/state_lora.py import torch import torch.nn as nn from peft.tuners.lora import LoraLayer class StateLoraLayer(LoraLayer): 一个简化的状态微调LoRA层实现。 它学习一个小的元参数向量用于生成LoRA权重。 def __init__( self, base_layer: nn.Module, adapter_name: str, r: int 8, # LoRA 秩 lora_alpha: int 32, lora_dropout: float 0.0, meta_dim: int 128, # 元参数的维度 **kwargs, ): super().__init__(base_layer, adapter_name) self.r r self.lora_alpha lora_alpha self.lora_dropout nn.Dropout(plora_dropout) if lora_dropout 0.0 else nn.Identity() # 传统的 LoRA 权重 A, B (被冻结仅作为生成目标的基础或初始值) self.lora_A nn.Parameter(torch.randn(base_layer.in_features, r), requires_gradFalse) self.lora_B nn.Parameter(torch.zeros(r, base_layer.out_features), requires_gradFalse) # 核心状态微调部分 # 1. 定义可学习的元参数 self.meta_params nn.Parameter(torch.randn(meta_dim)) # 2. 定义从元参数生成 LoRA 权重的轻量级网络 # 这里使用一个简单的两层MLP输入是元参数输出是展平的 A 和 B 的增量 total_lora_params base_layer.in_features * r r * base_layer.out_features self.meta_to_delta nn.Sequential( nn.Linear(meta_dim, 256), nn.ReLU(), nn.Linear(256, total_lora_params) # 输出维度等于 A 和 B 的总参数数 ) # 缩放因子 self.scaling lora_alpha / r def get_delta_weights(self): 根据当前元参数生成 LoRA 权重 A 和 B 的增量。 delta_flat self.meta_to_delta(self.meta_params) # 将扁平化的增量拆分为 A_delta 和 B_delta a_size self.base_layer.in_features * self.r a_delta_flat, b_delta_flat delta_flat[:a_size], delta_flat[a_size:] a_delta a_delta_flat.view(self.base_layer.in_features, self.r) b_delta b_delta_flat.view(self.r, self.base_layer.out_features) # 生成当前步骤的 LoRA 权重 current_A self.lora_A a_delta current_B self.lora_B b_delta return current_A, current_B def forward(self, x: torch.Tensor): previous_dtype x.dtype # 获取动态生成的 LoRA 权重 A, B self.get_delta_weights() # 执行 LoRA 前向传播 result self.base_layer(x) lora_output self.lora_dropout(x) A.to(x.dtype) B.to(x.dtype) result result lora_output * self.scaling result result.to(previous_dtype) return result4.3 配置与训练脚本接下来我们编写训练脚本集成状态微调层、并行控制策略和低显存技术。# configs/train_config.yaml model_name_or_path: Qwen/Qwen-7B-Chat # 基座模型 dataset_path: ./data/train.jsonl output_dir: ./outputs/state_lora_exp # LoRA 配置 lora_config: r: 16 lora_alpha: 32 lora_dropout: 0.1 target_modules: [q_proj, k_proj, v_proj, o_proj] # 针对 Qwen 的注意力模块 # 状态微调特定配置 use_state_tuning: true meta_dim: 256 # 训练参数 training_args: num_train_epochs: 3 per_device_train_batch_size: 2 # 小批次以适应低显存 gradient_accumulation_steps: 8 # 梯度累积模拟大批次 learning_rate: 2e-4 warmup_steps: 100 logging_steps: 10 save_steps: 200 fp16: true # 混合精度训练 gradient_checkpointing: true # 梯度检查点 optim: adamw_8bit # 使用8-bit Adam优化器bitsandbytes提供 # 并行任务配置 (模拟) parallel_tasks: - task_id: task_code lora_alpha: 64 learning_rate: 3e-4 - task_id: task_math lora_alpha: 32 learning_rate: 1e-4# scripts/train_state_parallel.py import os import yaml import torch from accelerate import Accelerator from transformers import ( AutoTokenizer, AutoModelForCausalLM, DataCollatorForSeq2Seq, TrainingArguments, Trainer ) from peft import get_peft_model, TaskType from datasets import load_dataset # 导入我们自定义的状态LoRA配置和层 from utils.state_lora import StateLoraConfig, StateLoraModel # 假设我们将上面的层封装成了PeftConfig和PeftModel def load_config(config_path): with open(config_path, r) as f: config yaml.safe_load(f) return config def main(): # 1. 加载配置 config load_config(./configs/train_config.yaml) accelerator Accelerator(cpu_offloadTrue) # 启用CPU Offloading # 2. 加载模型和分词器使用量化加载以节省显存 tokenizer AutoTokenizer.from_pretrained(config[model_name_or_path], trust_remote_codeTrue) tokenizer.pad_token tokenizer.eos_token # 设置填充token # 使用 bitsandbytes 进行 8-bit 量化加载 model AutoModelForCausalLM.from_pretrained( config[model_name_or_path], load_in_8bitTrue, # 关键8-bit量化 device_mapauto, # 自动分配模型层到GPU/CPU torch_dtypetorch.float16, trust_remote_codeTrue ) model.gradient_checkpointing_enable() # 启用梯度检查点 # 3. 准备数据集 dataset load_dataset(json, data_files{train: config[dataset_path]})[train] def tokenize_function(examples): # 简单分词实际应用应使用 apply_chat_template return tokenizer(examples[text], truncationTrue, paddingmax_length, max_length512) tokenized_dataset dataset.map(tokenize_function, batchedTrue, remove_columns[text]) # 4. 为每个并行任务创建并配置PEFT模型 peft_models [] training_args_list [] for i, task_config in enumerate(config[parallel_tasks]): print(f\n 准备并行任务 {i1}: {task_config[task_id]} ) # 创建状态LoRA配置 peft_config StateLoraConfig( task_typeTaskType.CAUSAL_LM, rconfig[lora_config][r], lora_alphatask_config.get(lora_alpha, config[lora_config][lora_alpha]), lora_dropoutconfig[lora_config][lora_dropout], target_modulesconfig[lora_config][target_modules], use_state_tuningconfig[lora_config][use_state_tuning], meta_dimconfig[lora_config][meta_dim] ) # 获取PEFT模型。注意我们为每个任务创建一个新的PEFT模型但它们共享底层的基础模型。 peft_model get_peft_model(model, peft_config) peft_model.print_trainable_parameters() # 打印可训练参数量 # 配置任务特定的训练参数 task_output_dir os.path.join(config[output_dir], task_config[task_id]) training_args TrainingArguments( output_dirtask_output_dir, num_train_epochsconfig[training_args][num_train_epochs], per_device_train_batch_sizeconfig[training_args][per_device_train_batch_size], gradient_accumulation_stepsconfig[training_args][gradient_accumulation_steps], learning_ratetask_config.get(learning_rate, config[training_args][learning_rate]), warmup_stepsconfig[training_args][warmup_steps], logging_stepsconfig[training_args][logging_steps], save_stepsconfig[training_args][save_steps], fp16config[training_args][fp16], gradient_checkpointingconfig[training_args][gradient_checkpointing], optimconfig[training_args][optim], report_tonone, # 禁用wandb等简化示例 ) peft_models.append(peft_model) training_args_list.append(training_args) # 5. 序列化执行并行任务模拟并行控制 data_collator DataCollatorForSeq2Seq(tokenizer, modelmodel, paddingTrue) for idx, (peft_model, args) in enumerate(zip(peft_models, training_args_list)): print(f\n 开始训练任务: {config[parallel_tasks][idx][task_id]}) # 使用 accelerate 准备当前任务的模型、优化器等 model, optimizer, train_dataloader accelerator.prepare(peft_model, ...) # 简化表示 trainer Trainer( modelpeft_model, argsargs, train_datasettokenized_dataset, data_collatordata_collator, tokenizertokenizer, ) trainer.train() # 任务训练完成后保存该任务的适配器权重即元参数 task_save_path os.path.join(args.output_dir, final_adapter) peft_model.save_pretrained(task_save_path) print(f任务适配器已保存至: {task_save_path}) # 关键步骤清理当前任务的计算图、释放显存为下一个任务做准备 # 将模型移出GPU清理缓存 peft_model.to(cpu) torch.cuda.empty_cache() # 注意基础模型仍在GPU中但已无计算图依赖。 print(\n所有并行任务训练完成) if __name__ __main__: main()4.4 运行与验证准备数据cd lora_advanced_tuning python scripts/data_prepare.py运行状态微调并行训练accelerate launch --num_processes1 scripts/train_state_parallel.py--num_processes1表示单GPU运行我们的并行控制是在单个GPU上序列化执行多个任务。验证与推理训练完成后每个任务会保存一个适配器。你可以加载基础模型和对应的适配器进行推理测试。from peft import PeftModel from transformers import AutoTokenizer, AutoModelForCausalLM, pipeline base_model AutoModelForCausalLM.from_pretrained(Qwen/Qwen-7B-Chat, load_in_8bitTrue, device_mapauto) tokenizer AutoTokenizer.from_pretrained(Qwen/Qwen-7B-Chat, trust_remote_codeTrue) # 加载任务一的适配器 task1_model PeftModel.from_pretrained(base_model, ./outputs/state_lora_exp/task_code/final_adapter) # 创建文本生成管道 pipe pipeline(text-generation, modeltask1_model, tokenizertokenizer) result pipe(写一个Python函数计算斐波那契数列。) print(result[0][generated_text])4.5 结果说明通过上述方案你可以观察到显存占用在单个消费级GPU如RTX 3090 24GB上同时管理两个微调任务成为可能。峰值显存主要由基础模型8-bit量化后约7-8GB、一个任务的激活/梯度约2-4GB和多个小型优化器状态组成。任务隔离每个任务的适配器即元参数是独立保存和加载的互不影响。灵活性通过调整meta_dim和meta_to_delta网络的结构可以在参数效率和表达能力之间进行权衡。5. 常见问题与排查思路在实现低显存LoRA方案时你可能会遇到以下问题问题现象可能原因排查思路与解决方案CUDA out of memory1. 批次大小过大。2. 梯度累积步数设置过小。3. 未启用梯度检查点或混合精度。4. 模型未成功量化load_in_8bit失效。5. 并行任务间显存未正确释放。1. 减小per_device_train_batch_size。2. 增大gradient_accumulation_steps。3. 确保gradient_checkpointingTrue和fp16True。4. 检查bitsandbytes版本和CUDA兼容性确保模型被正确量化查看model参数数据类型。5. 确保在切换任务时执行model.to(cpu)和torch.cuda.empty_cache()。训练损失不下降或为NaN1. 学习率过高。2. 混合精度训练不稳定。3. 状态微调中元参数生成网络meta_to_delta输出值过大。4. 数据预处理或分词错误。1. 降低学习率尝试1e-5到5e-5的范围。2. 尝试使用bf16如果硬件支持或暂时禁用fp16。3. 在meta_to_delta网络的输出层后添加Tanh或Sigmoid激活函数进行缩放或初始化权重更小。4. 检查tokenized_dataset的样本确保输入和标签格式正确。加载适配器后模型输出无变化1. 适配器未正确合并或激活。2. 状态微调适配器保存/加载的元参数或网络结构不匹配。3.target_modules设置错误未覆盖到关键层。1. 使用peft_model PeftModel.from_pretrained(base_model, adapter_path)后确保推理时调用的是peft_model。2. 检查保存的adapter_config.json和adapter_model.bin文件确保自定义的StateLoraConfig被正确保存和识别。加载时可能需要传入自定义的配置类。3. 确认基座模型的模块名称修正target_modules列表。对于Qwen通常是q_proj,k_proj,v_proj,o_proj。并行任务训练速度极慢1. CPU Offloading 过于频繁导致GPU-CPU数据传输成为瓶颈。2. 序列化执行总时间是各任务时间之和。1. 调整accelerate的offload_folder到更快的存储如SSD或减少卸载的数据量如只卸载优化器状态。2. 这是本方案为节省显存付出的代价。如果显存允许可考虑使用多GPU进行真正的数据并行。bitsandbytes相关错误1. CUDA版本不兼容。2. 安装的bitsandbytes版本不对。1. 确保CUDA版本与bitsandbytes预编译版本匹配。可能需要从源码编译bitsandbytes。2. 尝试pip install -U bitsandbytes或安装特定版本pip install bitsandbytes0.41.0。6. 最佳实践与工程建议将低显存LoRA方案应用于生产或严肃研究时请考虑以下建议渐进式复杂度不要一开始就使用最复杂的“状态微调”。首先用标准的LoRA权重微调配合梯度检查点、混合精度、4/8-bit量化这三板斧解决大部分显存问题。只有在多任务并行压力极大且标准方法仍不足时再引入状态微调和复杂的并行控制逻辑。量化策略选择训练阶段load_in_8bit(LLM.int8()) 通常足够且稳定。load_in_4bit(QLoRA) 能进一步节省显存但可能带来轻微的性能损失和更复杂的依赖。推理阶段可以使用GPTQ、AWQ等后训练量化方法获得更快的推理速度和更低的显存占用。监控与剖析使用nvidia-smi、torch.cuda.memory_allocated()或accelerate的trackers来监控显存使用情况。明确瓶颈是在模型参数、激活、梯度还是优化器状态。状态微调的设计本文的StateLoraLayer是一个概念演示。在实际应用中meta_to_delta网络的设计至关重要。可以考虑更高效的参数化如使用超网络HyperNetwork、低秩分解或条件生成。任务条件化让元参数θ也接收任务ID或任务描述作为输入实现更灵活的多任务学习。共享与独享部分元参数可以在不同任务间共享部分保持独立以平衡容量和效率。保存与部署状态微调保存的是元参数和生成网络。部署时需要同时加载基础模型、生成网络和元参数。考虑将生成网络和元参数打包成一个独立的、轻量的推理模块。对于标准LoRA训练完成后可以使用merge_and_unload()将适配器权重合并到基础模型中导出为一个完整的模型文件简化部署。安全与测试环境隔离在Docker容器或虚拟环境中进行实验确保依赖库版本一致。小规模验证先用1%的数据和1个epoch跑通全流程确认代码无误、显存可控后再开始大规模训练。备份检查点定期保存训练检查点并验证检查点可以成功加载和恢复训练。从权重微调到状态微调的演进代表了参数高效微调技术向更深层次资源优化的探索。并行控制下的低显存方案不是银弹而是一套组合拳需要根据你的具体硬件条件、任务数量和模型规模灵活搭配。核心思路始终是共享一切可共享的基础模型量化一切可量化的模型权重动态生成一切可生成的适配器权重并妥善管理生命周期显存复用。掌握这套方案后你可以更从容地在大模型上进行多任务学习、超参数搜索和个性化定制。建议你从修改示例代码开始尝试调整meta_dim、设计不同的meta_to_delta网络结构并在你自己的数据集上验证效果。实践过程中遇到的挑战和解决方案将是你在AIGC工程化道路上最宝贵的经验。
返回列表