Buckets:

download
raw
16.6 kB
import"../chunks/DsnmJJEf.js";import{i as B,h as k,C as E,H as t,a as e,b as n,E as F,s as S}from"../chunks/DdZvggmf.js";import{p as A,o as I,s as a,f as N,a as i,b as Q,c as p,n as z}from"../chunks/BbekZcyp.js";import{H as c}from"../chunks/BcnRgdDK.js";const q='{"title":"文生图","local":"文生图","sections":[{"title":"脚本参数","local":"脚本参数","sections":[{"title":"Min-SNR加权策略","local":"min-snr加权策略","sections":[],"depth":3}],"depth":2},{"title":"训练脚本解析","local":"训练脚本解析","sections":[],"depth":2},{"title":"启动脚本","local":"启动脚本","sections":[],"depth":2},{"title":"后续步骤","local":"后续步骤","sections":[],"depth":2}],"depth":1}';var L=p('<meta name="hf:doc:metadata"/>'),H=p('<p>以 <a href="https://huggingface.co/datasets/lambdalabs/naruto-blip-captions" rel="nofollow">火影忍者BLIP标注数据集</a> 为例训练生成火影角色。设置环境变量 <code>MODEL_NAME</code> 和 <code>dataset_name</code> 指定模型和数据集(Hub或本地路径)。多GPU训练需在 <code>accelerate launch</code> 命令中添加 <code>--multi_gpu</code> 参数。</p> <blockquote class="tip"><p>使用本地数据集时,设置 <code>TRAIN_DIR</code> 和 <code>OUTPUT_DIR</code> 环境变量为数据集路径和模型保存路径。</p></blockquote> <!>',1),P=p('<p></p> <!> <!> <blockquote class="warning"><p>文生图训练脚本目前处于实验阶段,容易出现过拟合和灾难性遗忘等问题。建议尝试不同超参数以获得最佳数据集适配效果。</p></blockquote> <p>Stable Diffusion 等文生图模型能够根据文本提示生成对应图像。</p> <p>模型训练对硬件要求较高,但启用 <code>gradient_checkpointing</code> 和 <code>mixed_precision</code> 后,可在单块24GB显存GPU上完成训练。如需更大批次或更快训练速度,建议使用30GB以上显存的GPU设备。通过启用 <a href="../optimization/xformers">xFormers</a> 内存高效注意力机制可降低显存占用。</p> <p>本指南将详解 <a href="https://github.com/huggingface/diffusers/blob/main/examples/text_to_image/train_text_to_image.py" rel="nofollow">train_text_to_image.py</a> 训练脚本,助您掌握其原理并适配自定义需求。</p> <p>运行脚本前请确保已从源码安装库:</p> <!> <p>然后进入包含训练脚本的示例目录,安装对应依赖:</p> <!> <blockquote class="tip"><p>🤗 Accelerate 是支持多GPU/TPU训练和混合精度的工具库,能根据硬件环境自动配置训练参数。参阅 🤗 Accelerate <a href="https://huggingface.co/docs/accelerate/quicktour" rel="nofollow">快速入门</a> 了解更多。</p></blockquote> <p>初始化 🤗 Accelerate 环境:</p> <!> <p>要创建默认配置环境(不进行交互式选择):</p> <!> <p>若环境不支持交互式shell(如notebook),可使用:</p> <!> <p>最后,如需在自定义数据集上训练,请参阅 <a href="create_dataset">创建训练数据集</a> 指南了解如何准备适配脚本的数据集。</p> <!> <blockquote class="tip"><p>以下重点介绍脚本中影响训练效果的关键参数,如需完整参数说明可查阅 <a href="https://github.com/huggingface/diffusers/blob/main/examples/text_to_image/train_text_to_image.py" rel="nofollow">脚本源码</a>。如有疑问欢迎反馈。</p></blockquote> <p>训练脚本提供丰富参数供自定义训练流程,所有参数及说明详见 <a href="https://github.com/huggingface/diffusers/blob/8959c5b9dec1c94d6ba482c94a58d2215c5fd026/examples/text_to_image/train_text_to_image.py#L193" rel="nofollow"><code>parse_args()</code></a> 函数。该函数为每个参数提供默认值(如批次大小、学习率等),也可通过命令行参数覆盖。</p> <p>例如使用fp16混合精度加速训练:</p> <!> <p>基础重要参数包括:</p> <ul><li><code>--pretrained_model_name_or_path</code>: Hub模型名称或本地预训练模型路径</li> <li><code>--dataset_name</code>: Hub数据集名称或本地训练数据集路径</li> <li><code>--image_column</code>: 数据集中图像列名</li> <li><code>--caption_column</code>: 数据集中文本列名</li> <li><code>--output_dir</code>: 模型保存路径</li> <li><code>--push_to_hub</code>: 是否将训练模型推送至Hub</li> <li><code>--checkpointing_steps</code>: 模型检查点保存步数;训练中断时可添加 <code>--resume_from_checkpoint</code> 从该检查点恢复训练</li></ul> <!> <p><a href="https://huggingface.co/papers/2303.09556" rel="nofollow">Min-SNR</a> 加权策略通过重新平衡损失函数加速模型收敛。训练脚本支持预测 <code>epsilon</code>(噪声)或 <code>v_prediction</code>,而Min-SNR兼容两种预测类型。</p> <p>添加 <code>--snr_gamma</code> 参数并设为推荐值5.0:</p> <!> <p>可通过此 <a href="https://wandb.ai/sayakpaul/text2image-finetune-minsnr" rel="nofollow">Weights and Biases</a> 报告比较不同 <code>snr_gamma</code> 值的损失曲面。小数据集上Min-SNR效果可能不如大数据集显著。</p> <!> <p>数据集预处理代码和训练循环位于 <a href="https://github.com/huggingface/diffusers/blob/8959c5b9dec1c94d6ba482c94a58d2215c5fd026/examples/text_to_image/train_text_to_image.py#L490" rel="nofollow"><code>main()</code></a> 函数,自定义修改需在此处进行。</p> <p><code>train_text_to_image</code> 脚本首先 <a href="https://github.com/huggingface/diffusers/blob/8959c5b9dec1c94d6ba482c94a58d2215c5fd026/examples/text_to_image/train_text_to_image.py#L543" rel="nofollow">加载调度器</a> 和分词器,此处可替换其他调度器:</p> <!> <p>接着 <a href="https://github.com/huggingface/diffusers/blob/8959c5b9dec1c94d6ba482c94a58d2215c5fd026/examples/text_to_image/train_text_to_image.py#L619" rel="nofollow">加载UNet模型</a>:</p> <!> <p>随后对数据集的文本和图像列进行预处理。<a href="https://github.com/huggingface/diffusers/blob/8959c5b9dec1c94d6ba482c94a58d2215c5fd026/examples/text_to_image/train_text_to_image.py#L724" rel="nofollow"><code>tokenize_captions</code></a> 函数处理文本分词,<a href="https://github.com/huggingface/diffusers/blob/8959c5b9dec1c94d6ba482c94a58d2215c5fd026/examples/text_to_image/train_text_to_image.py#L742" rel="nofollow"><code>train_transforms</code></a> 定义图像增强策略,二者集成于 <code>preprocess_train</code>:</p> <!> <p>最后,<a href="https://github.com/huggingface/diffusers/blob/8959c5b9dec1c94d6ba482c94a58d2215c5fd026/examples/text_to_image/train_text_to_image.py#L878" rel="nofollow">训练循环</a> 处理剩余流程:图像编码为潜空间、添加噪声、计算文本嵌入条件、更新模型参数、保存并推送模型至Hub。想深入了解训练循环原理,可参阅 <a href="../using-diffusers/write_own_pipeline">理解管道、模型与调度器</a> 教程,该教程解析了去噪过程的核心逻辑。</p> <!> <p>完成所有配置后,即可启动训练脚本!🚀</p> <!> <p>训练完成后,即可使用新模型进行推理:</p> <!> <!> <p>恭喜完成文生图模型训练!如需进一步使用模型,以下指南可能有所帮助:</p> <ul><li>了解如何加载 <a href="../using-diffusers/loading_adapters#LoRA">LoRA权重</a> 进行推理(如果训练时使用了LoRA)</li> <li>在 <a href="../using-diffusers/conditional_image_generation">文生图</a> 任务指南中,了解引导尺度等参数或提示词加权等技术如何控制生成效果</li></ul> <!> <p></p>',1);function aa(V,v){A(v,!1),I(()=>{new URLSearchParams(window.location.search).get("fw")}),B();var d=P();k("1wa9aw0",l=>{var s=L();S(s,"content",q),i(l,s)});var r=a(N(d),2);E(r,{containerStyle:"float: right; margin-left: 10px; display: inline-flex; position: relative; z-index: 10;"});var h=a(r,2);t(h,{title:"文生图",local:"文生图",headingTag:"h1"});var m=a(h,12);e(m,{code:"Z2l0JTIwY2xvbmUlMjBodHRwcyUzQSUyRiUyRmdpdGh1Yi5jb20lMkZodWdnaW5nZmFjZSUyRmRpZmZ1c2VycyUwQWNkJTIwZGlmZnVzZXJzJTBBcGlwJTIwaW5zdGFsbCUyMC4=",highlighted:`git <span class="hljs-built_in">clone</span> https://github.com/huggingface/diffusers
<span class="hljs-built_in">cd</span> diffusers
pip install .`,lang:"bash",wrap:!1});var u=a(m,4);n(u,{id:"installation",options:["PyTorch"],children:(l,s)=>{c(l,{id:"installation",option:"PyTorch",children:(o,x)=>{e(o,{code:"Y2QlMjBleGFtcGxlcyUyRnRleHRfdG9faW1hZ2UlMEFwaXAlMjBpbnN0YWxsJTIwLXIlMjByZXF1aXJlbWVudHMudHh0",highlighted:`<span class="hljs-built_in">cd</span> examples/text_to_image
pip install -r requirements.txt`,lang:"bash",wrap:!1})},$$slots:{default:!0}})},$$slots:{default:!0}});var g=a(u,6);e(g,{code:"YWNjZWxlcmF0ZSUyMGNvbmZpZw==",highlighted:"accelerate config",lang:"bash",wrap:!1});var b=a(g,4);e(b,{code:"YWNjZWxlcmF0ZSUyMGNvbmZpZyUyMGRlZmF1bHQ=",highlighted:"accelerate config default",lang:"bash",wrap:!1});var M=a(b,4);e(M,{code:"ZnJvbSUyMGFjY2VsZXJhdGUudXRpbHMlMjBpbXBvcnQlMjB3cml0ZV9iYXNpY19jb25maWclMEElMEF3cml0ZV9iYXNpY19jb25maWcoKQ==",highlighted:`<span class="hljs-keyword">from</span> accelerate.utils <span class="hljs-keyword">import</span> write_basic_config
write_basic_config()`,lang:"py",wrap:!1});var f=a(M,4);t(f,{title:"脚本参数",local:"脚本参数",headingTag:"h2"});var _=a(f,8);e(_,{code:"YWNjZWxlcmF0ZSUyMGxhdW5jaCUyMHRyYWluX3RleHRfdG9faW1hZ2UucHklMjAlNUMlMEElMjAlMjAtLW1peGVkX3ByZWNpc2lvbiUzRCUyMmZwMTYlMjI=",highlighted:`accelerate launch train_text_to_image.py \\
--mixed_precision=<span class="hljs-string">&quot;fp16&quot;</span>`,lang:"bash",wrap:!1});var y=a(_,6);t(y,{title:"Min-SNR加权策略",local:"min-snr加权策略",headingTag:"h3"});var U=a(y,6);e(U,{code:"YWNjZWxlcmF0ZSUyMGxhdW5jaCUyMHRyYWluX3RleHRfdG9faW1hZ2UucHklMjAlNUMlMEElMjAlMjAtLXNucl9nYW1tYSUzRDUuMA==",highlighted:`accelerate launch train_text_to_image.py \\
--snr_gamma=5.0`,lang:"bash",wrap:!1});var Z=a(U,4);t(Z,{title:"训练脚本解析",local:"训练脚本解析",headingTag:"h2"});var j=a(Z,6);e(j,{code:"bm9pc2Vfc2NoZWR1bGVyJTIwJTNEJTIwRERQTVNjaGVkdWxlci5mcm9tX3ByZXRyYWluZWQoYXJncy5wcmV0cmFpbmVkX21vZGVsX25hbWVfb3JfcGF0aCUyQyUyMHN1YmZvbGRlciUzRCUyMnNjaGVkdWxlciUyMiklMEF0b2tlbml6ZXIlMjAlM0QlMjBDTElQVG9rZW5pemVyLmZyb21fcHJldHJhaW5lZCglMEElMjAlMjAlMjAlMjBhcmdzLnByZXRyYWluZWRfbW9kZWxfbmFtZV9vcl9wYXRoJTJDJTIwc3ViZm9sZGVyJTNEJTIydG9rZW5pemVyJTIyJTJDJTIwcmV2aXNpb24lM0RhcmdzLnJldmlzaW9uJTBBKQ==",highlighted:`noise_scheduler = DDPMScheduler.from_pretrained(args.pretrained_model_name_or_path, subfolder=<span class="hljs-string">&quot;scheduler&quot;</span>)
tokenizer = CLIPTokenizer.from_pretrained(
args.pretrained_model_name_or_path, subfolder=<span class="hljs-string">&quot;tokenizer&quot;</span>, revision=args.revision
)`,lang:"py",wrap:!1});var w=a(j,4);e(w,{code:"bG9hZF9tb2RlbCUyMCUzRCUyMFVOZXQyRENvbmRpdGlvbk1vZGVsLmZyb21fcHJldHJhaW5lZChpbnB1dF9kaXIlMkMlMjBzdWJmb2xkZXIlM0QlMjJ1bmV0JTIyKSUwQW1vZGVsLnJlZ2lzdGVyX3RvX2NvbmZpZygqKmxvYWRfbW9kZWwuY29uZmlnKSUwQSUwQW1vZGVsLmxvYWRfc3RhdGVfZGljdChsb2FkX21vZGVsLnN0YXRlX2RpY3QoKSk=",highlighted:`load_model = UNet2DConditionModel.from_pretrained(input_dir, subfolder=<span class="hljs-string">&quot;unet&quot;</span>)
model.register_to_config(**load_model.config)
model.load_state_dict(load_model.state_dict())`,lang:"py",wrap:!1});var W=a(w,4);e(W,{code:"ZGVmJTIwcHJlcHJvY2Vzc190cmFpbihleGFtcGxlcyklM0ElMEElMjAlMjAlMjAlMjBpbWFnZXMlMjAlM0QlMjAlNUJpbWFnZS5jb252ZXJ0KCUyMlJHQiUyMiklMjBmb3IlMjBpbWFnZSUyMGluJTIwZXhhbXBsZXMlNUJpbWFnZV9jb2x1bW4lNUQlNUQlMEElMjAlMjAlMjAlMjBleGFtcGxlcyU1QiUyMnBpeGVsX3ZhbHVlcyUyMiU1RCUyMCUzRCUyMCU1QnRyYWluX3RyYW5zZm9ybXMoaW1hZ2UpJTIwZm9yJTIwaW1hZ2UlMjBpbiUyMGltYWdlcyU1RCUwQSUyMCUyMCUyMCUyMGV4YW1wbGVzJTVCJTIyaW5wdXRfaWRzJTIyJTVEJTIwJTNEJTIwdG9rZW5pemVfY2FwdGlvbnMoZXhhbXBsZXMpJTBBJTIwJTIwJTIwJTIwcmV0dXJuJTIwZXhhbXBsZXM=",highlighted:`<span class="hljs-keyword">def</span> <span class="hljs-title function_">preprocess_train</span>(<span class="hljs-params">examples</span>):
images = [image.convert(<span class="hljs-string">&quot;RGB&quot;</span>) <span class="hljs-keyword">for</span> image <span class="hljs-keyword">in</span> examples[image_column]]
examples[<span class="hljs-string">&quot;pixel_values&quot;</span>] = [train_transforms(image) <span class="hljs-keyword">for</span> image <span class="hljs-keyword">in</span> images]
examples[<span class="hljs-string">&quot;input_ids&quot;</span>] = tokenize_captions(examples)
<span class="hljs-keyword">return</span> examples`,lang:"py",wrap:!1});var J=a(W,4);t(J,{title:"启动脚本",local:"启动脚本",headingTag:"h2"});var R=a(J,4);n(R,{id:"training-inference",options:["PyTorch"],children:(l,s)=>{c(l,{id:"training-inference",option:"PyTorch",children:(o,x)=>{var T=H(),Y=a(N(T),4);e(Y,{code:"ZXhwb3J0JTIwTU9ERUxfTkFNRSUzRCUyMnN0YWJsZS1kaWZmdXNpb24tdjEtNSUyRnN0YWJsZS1kaWZmdXNpb24tdjEtNSUyMiUwQWV4cG9ydCUyMGRhdGFzZXRfbmFtZSUzRCUyMmxhbWJkYWxhYnMlMkZuYXJ1dG8tYmxpcC1jYXB0aW9ucyUyMiUwQSUwQWFjY2VsZXJhdGUlMjBsYXVuY2glMjAtLW1peGVkX3ByZWNpc2lvbiUzRCUyMmZwMTYlMjIlMjAlMjB0cmFpbl90ZXh0X3RvX2ltYWdlLnB5JTIwJTVDJTBBJTIwJTIwLS1wcmV0cmFpbmVkX21vZGVsX25hbWVfb3JfcGF0aCUzRCUyNE1PREVMX05BTUUlMjAlNUMlMEElMjAlMjAtLWRhdGFzZXRfbmFtZSUzRCUyNGRhdGFzZXRfbmFtZSUyMCU1QyUwQSUyMCUyMC0tdXNlX2VtYSUyMCU1QyUwQSUyMCUyMC0tcmVzb2x1dGlvbiUzRDUxMiUyMC0tY2VudGVyX2Nyb3AlMjAtLXJhbmRvbV9mbGlwJTIwJTVDJTBBJTIwJTIwLS10cmFpbl9iYXRjaF9zaXplJTNEMSUyMCU1QyUwQSUyMCUyMC0tZ3JhZGllbnRfYWNjdW11bGF0aW9uX3N0ZXBzJTNENCUyMCU1QyUwQSUyMCUyMC0tZ3JhZGllbnRfY2hlY2twb2ludGluZyUyMCU1QyUwQSUyMCUyMC0tbWF4X3RyYWluX3N0ZXBzJTNEMTUwMDAlMjAlNUMlMEElMjAlMjAtLWxlYXJuaW5nX3JhdGUlM0QxZS0wNSUyMCU1QyUwQSUyMCUyMC0tbWF4X2dyYWRfbm9ybSUzRDElMjAlNUMlMEElMjAlMjAtLWVuYWJsZV94Zm9ybWVyc19tZW1vcnlfZWZmaWNpZW50X2F0dGVudGlvbiUyMCU1QyUwQSUyMCUyMC0tbHJfc2NoZWR1bGVyJTNEJTIyY29uc3RhbnQlMjIlMjAtLWxyX3dhcm11cF9zdGVwcyUzRDAlMjAlNUMlMEElMjAlMjAtLW91dHB1dF9kaXIlM0QlMjJzZC1uYXJ1dG8tbW9kZWwlMjIlMjAlNUMlMEElMjAlMjAtLXB1c2hfdG9faHVi",highlighted:`<span class="hljs-built_in">export</span> MODEL_NAME=<span class="hljs-string">&quot;stable-diffusion-v1-5/stable-diffusion-v1-5&quot;</span>
<span class="hljs-built_in">export</span> dataset_name=<span class="hljs-string">&quot;lambdalabs/naruto-blip-captions&quot;</span>
accelerate launch --mixed_precision=<span class="hljs-string">&quot;fp16&quot;</span> train_text_to_image.py \\
--pretrained_model_name_or_path=<span class="hljs-variable">$MODEL_NAME</span> \\
--dataset_name=<span class="hljs-variable">$dataset_name</span> \\
--use_ema \\
--resolution=512 --center_crop --random_flip \\
--train_batch_size=1 \\
--gradient_accumulation_steps=4 \\
--gradient_checkpointing \\
--max_train_steps=15000 \\
--learning_rate=1e-05 \\
--max_grad_norm=1 \\
--enable_xformers_memory_efficient_attention \\
--lr_scheduler=<span class="hljs-string">&quot;constant&quot;</span> --lr_warmup_steps=0 \\
--output_dir=<span class="hljs-string">&quot;sd-naruto-model&quot;</span> \\
--push_to_hub`,lang:"bash",wrap:!1}),i(o,T)},$$slots:{default:!0}})},$$slots:{default:!0}});var G=a(R,4);n(G,{id:"training-inference",options:["PyTorch"],children:(l,s)=>{c(l,{id:"training-inference",option:"PyTorch",children:(o,x)=>{e(o,{code:"ZnJvbSUyMGRpZmZ1c2VycyUyMGltcG9ydCUyMFN0YWJsZURpZmZ1c2lvblBpcGVsaW5lJTBBaW1wb3J0JTIwdG9yY2glMEElMEFwaXBlbGluZSUyMCUzRCUyMFN0YWJsZURpZmZ1c2lvblBpcGVsaW5lLmZyb21fcHJldHJhaW5lZCglMjJwYXRoJTJGdG8lMkZzYXZlZF9tb2RlbCUyMiUyQyUyMGR0eXBlJTNEdG9yY2guZmxvYXQxNiUyQyUyMHVzZV9zYWZldGVuc29ycyUzRFRydWUpLnRvKCUyMmN1ZGElMjIpJTBBJTBBaW1hZ2UlMjAlM0QlMjBwaXBlbGluZShwcm9tcHQlM0QlMjJ5b2RhJTIyKS5pbWFnZXMlNUIwJTVEJTBBaW1hZ2Uuc2F2ZSglMjJ5b2RhLW5hcnV0by5wbmclMjIp",highlighted:`<span class="hljs-keyword">from</span> diffusers <span class="hljs-keyword">import</span> StableDiffusionPipeline
<span class="hljs-keyword">import</span> torch
pipeline = StableDiffusionPipeline.from_pretrained(<span class="hljs-string">&quot;path/to/saved_model&quot;</span>, dtype=torch.float16, use_safetensors=<span class="hljs-literal">True</span>).to(<span class="hljs-string">&quot;cuda&quot;</span>)
image = pipeline(prompt=<span class="hljs-string">&quot;yoda&quot;</span>).images[<span class="hljs-number">0</span>]
image.save(<span class="hljs-string">&quot;yoda-naruto.png&quot;</span>)`,lang:"py",wrap:!1})},$$slots:{default:!0}})},$$slots:{default:!0}});var X=a(G,2);t(X,{title:"后续步骤",local:"后续步骤",headingTag:"h2"});var C=a(X,6);F(C,{source:"https://github.com/huggingface/diffusers/blob/main/docs/source/zh/training/text2image.md"}),z(2),i(V,d),Q()}export{aa as component};

Xet Storage Details

Size:
16.6 kB
·
Xet hash:
1f7fc9de06605f864edff2f0aa6e72323c6bc7aef8d679ed49481460a5337920

Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.