Buckets:

download
raw
11.8 kB
import"../chunks/DsnmJJEf.js";import{i as Z,h as v,C as k,H as a,a as l,E as _,s as C}from"../chunks/CmJXCtRL.js";import{p as I,o as V,s,f as W,a as T,b as B,c as N,d as U,r as X,n as R}from"../chunks/DK803DsY.js";const E='{"title":"AutoModel","local":"automodel","sections":[{"title":"Custom models","local":"custom-models","sections":[{"title":"Saving custom models","local":"saving-custom-models","sections":[],"depth":3}],"depth":2}],"depth":1}';var G=U('<meta name="hf:doc:metadata"/>'),A=U('<p></p> <!> <!> <p>The <a href="/docs/diffusers/pr_14313/en/api/models/auto_model#diffusers.AutoModel">AutoModel</a> class automatically detects and loads the correct model class (UNet, transformer, VAE) from a <code>config.json</code> file. You don’t need to know the specific model class name ahead of time. It supports data types and device placement, and works across model types and libraries.</p> <p>The example below loads a transformer from Diffusers and a text encoder from Transformers. Use the <code>subfolder</code> parameter to specify where to load the <code>config.json</code> file from.</p> <!> <!> <p><a href="/docs/diffusers/pr_14313/en/api/models/auto_model#diffusers.AutoModel">AutoModel</a> also loads models from the <a href="https://huggingface.co/models" rel="nofollow">Hub</a> that aren’t included in Diffusers. Set <code>trust_remote_code=True</code> in <a href="/docs/diffusers/pr_14313/en/api/models/auto_model#diffusers.AutoModel.from_pretrained">AutoModel.from_pretrained()</a> to load custom models.</p> <p>A custom model repository needs a Python module with the model class, and a <code>config.json</code> with an <code>auto_map</code> entry that maps <code>"AutoModel"</code> to <code>"module_file.ClassName"</code>.</p> <!> <p>The <code>config.json</code> includes the <code>auto_map</code> field pointing to the custom class.</p> <!> <p>Then load it with <code>trust_remote_code=True</code>.</p> <!> <p>For a real-world example, <a href="https://huggingface.co/Overworld/Waypoint-1-Small/tree/main/transformer" rel="nofollow">Overworld/Waypoint-1-Small</a> hosts a custom <code>WorldModel</code> class across several modules in its <code>transformer</code> subfolder.</p> <!> <!> <p>If the custom model inherits from the <a href="/docs/diffusers/pr_14313/en/api/models/overview#diffusers.ModelMixin">ModelMixin</a> class, it gets access to the same features as Diffusers model classes, like <a href="../optimization/fp16#regional-compilation">regional compilation</a> and <a href="../optimization/memory#group-offloading">group offloading</a>.</p> <blockquote class="warning"><p>As a precaution with <code>trust_remote_code=True</code>, pass a commit hash to the <code>revision</code> argument in <a href="/docs/diffusers/pr_14313/en/api/models/auto_model#diffusers.AutoModel.from_pretrained">AutoModel.from_pretrained()</a> to make sure the code hasn’t been updated with new malicious code (unless you fully trust the model owners).</p> <!></blockquote> <!> <p>Use <code>register_for_auto_class()</code> to add the <code>auto_map</code> entry to <code>config.json</code> automatically when saving. This avoids having to manually edit the config file.</p> <!> <p>The saved <code>config.json</code> will include the <code>auto_map</code> field.</p> <!> <blockquote class="note"><p>Learn more about implementing custom models in the <a href="../using-diffusers/custom_pipeline_overview#community-components">Community components</a> guide.</p></blockquote> <!> <p></p>',1);function Y(w,j){I(j,!1),V(()=>{new URLSearchParams(window.location.search).get("fw")}),Z();var e=A();v("1vmbd58",J=>{var f=G();C(f,"content",E),T(J,f)});var t=s(W(e),2);k(t,{containerStyle:"float: right; margin-left: 10px; display: inline-flex; position: relative; z-index: 10;"});var n=s(t,2);a(n,{title:"AutoModel",local:"automodel",headingTag:"h1"});var d=s(n,6);l(d,{code:"aW1wb3J0JTIwdG9yY2glMEFmcm9tJTIwZGlmZnVzZXJzJTIwaW1wb3J0JTIwQXV0b01vZGVsJTJDJTIwRGlmZnVzaW9uUGlwZWxpbmUlMEElMEF0cmFuc2Zvcm1lciUyMCUzRCUyMEF1dG9Nb2RlbC5mcm9tX3ByZXRyYWluZWQoJTBBJTIwJTIwJTIwJTIwJTIyUXdlbiUyRlF3ZW4tSW1hZ2UlMjIlMkMlMjBzdWJmb2xkZXIlM0QlMjJ0cmFuc2Zvcm1lciUyMiUyQyUyMGR0eXBlJTNEdG9yY2guYmZsb2F0MTYlMkMlMjBkZXZpY2VfbWFwJTNEJTIyY3VkYSUyMiUwQSklMEElMEF0ZXh0X2VuY29kZXIlMjAlM0QlMjBBdXRvTW9kZWwuZnJvbV9wcmV0cmFpbmVkKCUwQSUyMCUyMCUyMCUyMCUyMlF3ZW4lMkZRd2VuLUltYWdlJTIyJTJDJTIwc3ViZm9sZGVyJTNEJTIydGV4dF9lbmNvZGVyJTIyJTJDJTIwZHR5cGUlM0R0b3JjaC5iZmxvYXQxNiUyQyUyMGRldmljZV9tYXAlM0QlMjJjdWRhJTIyJTBBKQ==",highlighted:`<span class="hljs-keyword">import</span> torch
<span class="hljs-keyword">from</span> diffusers <span class="hljs-keyword">import</span> AutoModel, DiffusionPipeline
transformer = AutoModel.from_pretrained(
<span class="hljs-string">&quot;Qwen/Qwen-Image&quot;</span>, subfolder=<span class="hljs-string">&quot;transformer&quot;</span>, dtype=torch.bfloat16, device_map=<span class="hljs-string">&quot;cuda&quot;</span>
)
text_encoder = AutoModel.from_pretrained(
<span class="hljs-string">&quot;Qwen/Qwen-Image&quot;</span>, subfolder=<span class="hljs-string">&quot;text_encoder&quot;</span>, dtype=torch.bfloat16, device_map=<span class="hljs-string">&quot;cuda&quot;</span>
)`,lang:"py",wrap:!1});var r=s(d,2);a(r,{title:"Custom models",local:"custom-models",headingTag:"h2"});var c=s(r,6);l(c,{code:"Y3VzdG9tJTJGY3VzdG9tLXRyYW5zZm9ybWVyLW1vZGVsJTJGJTBBJUUyJTk0JTlDJUUyJTk0JTgwJUUyJTk0JTgwJTIwY29uZmlnLmpzb24lMEElRTIlOTQlOUMlRTIlOTQlODAlRTIlOTQlODAlMjBteV9tb2RlbC5weSUwQSVFMiU5NCU5NCVFMiU5NCU4MCVFMiU5NCU4MCUyMGRpZmZ1c2lvbl9weXRvcmNoX21vZGVsLnNhZmV0ZW5zb3Jz",highlighted:`<span class="hljs-literal">custom</span>/<span class="hljs-literal">custom</span>-transformer-model/
├── config.json
├── my_model.py
└── diffusion_pytorch_model.safetensors`,lang:"",wrap:!1});var i=s(c,4);l(i,{code:"JTdCJTBBJTIwJTIwJTIyYXV0b19tYXAlMjIlM0ElMjAlN0IlMEElMjAlMjAlMjAlMjAlMjJBdXRvTW9kZWwlMjIlM0ElMjAlMjJteV9tb2RlbC5NeUN1c3RvbU1vZGVsJTIyJTBBJTIwJTIwJTdEJTBBJTdE",highlighted:`<span class="hljs-punctuation">{</span>
<span class="hljs-attr">&quot;auto_map&quot;</span><span class="hljs-punctuation">:</span> <span class="hljs-punctuation">{</span>
<span class="hljs-attr">&quot;AutoModel&quot;</span><span class="hljs-punctuation">:</span> <span class="hljs-string">&quot;my_model.MyCustomModel&quot;</span>
<span class="hljs-punctuation">}</span>
<span class="hljs-punctuation">}</span>`,lang:"json",wrap:!1});var p=s(i,4);l(p,{code:"aW1wb3J0JTIwdG9yY2glMEFmcm9tJTIwZGlmZnVzZXJzJTIwaW1wb3J0JTIwQXV0b01vZGVsJTBBJTBBdHJhbnNmb3JtZXIlMjAlM0QlMjBBdXRvTW9kZWwuZnJvbV9wcmV0cmFpbmVkKCUwQSUyMCUyMCUyMCUyMCUyMmN1c3RvbSUyRmN1c3RvbS10cmFuc2Zvcm1lci1tb2RlbCUyMiUyQyUyMHRydXN0X3JlbW90ZV9jb2RlJTNEVHJ1ZSUyQyUyMGR0eXBlJTNEdG9yY2guYmZsb2F0MTYlMkMlMjBkZXZpY2VfbWFwJTNEJTIyY3VkYSUyMiUwQSk=",highlighted:`<span class="hljs-keyword">import</span> torch
<span class="hljs-keyword">from</span> diffusers <span class="hljs-keyword">import</span> AutoModel
transformer = AutoModel.from_pretrained(
<span class="hljs-string">&quot;custom/custom-transformer-model&quot;</span>, trust_remote_code=<span class="hljs-literal">True</span>, dtype=torch.bfloat16, device_map=<span class="hljs-string">&quot;cuda&quot;</span>
)`,lang:"py",wrap:!1});var m=s(p,4);l(m,{code:"dHJhbnNmb3JtZXIlMkYlMEElRTIlOTQlOUMlRTIlOTQlODAlRTIlOTQlODAlMjBjb25maWcuanNvbiUyMCUyMCUyMCUyMCUyMCUyMCUyMCUyMCUyMCUyMCUyMyUyMGF1dG9fbWFwJTNBJTIwJTIybW9kZWwuV29ybGRNb2RlbCUyMiUwQSVFMiU5NCU5QyVFMiU5NCU4MCVFMiU5NCU4MCUyMG1vZGVsLnB5JTBBJUUyJTk0JTlDJUUyJTk0JTgwJUUyJTk0JTgwJTIwYXR0bi5weSUwQSVFMiU5NCU5QyVFMiU5NCU4MCVFMiU5NCU4MCUyMG5uLnB5JTBBJUUyJTk0JTlDJUUyJTk0JTgwJUUyJTk0JTgwJTIwY2FjaGUucHklMEElRTIlOTQlOUMlRTIlOTQlODAlRTIlOTQlODAlMjBxdWFudGl6ZS5weSUwQSVFMiU5NCU5QyVFMiU5NCU4MCVFMiU5NCU4MCUyMF9faW5pdF9fLnB5JTBBJUUyJTk0JTk0JUUyJTk0JTgwJUUyJTk0JTgwJTIwZGlmZnVzaW9uX3B5dG9yY2hfbW9kZWwuc2FmZXRlbnNvcnM=",highlighted:`transformer/
├── config.json # auto_map: <span class="hljs-string">&quot;model.WorldModel&quot;</span>
├── model.<span class="hljs-keyword">py</span>
├── attn.<span class="hljs-keyword">py</span>
├── <span class="hljs-keyword">nn</span>.<span class="hljs-keyword">py</span>
├── cache.<span class="hljs-keyword">py</span>
├── quantize.<span class="hljs-keyword">py</span>
├── __init__.<span class="hljs-keyword">py</span>
└── diffusion_pytorch_model.safetensors`,lang:"",wrap:!1});var u=s(m,2);l(u,{code:"aW1wb3J0JTIwdG9yY2glMEFmcm9tJTIwZGlmZnVzZXJzJTIwaW1wb3J0JTIwQXV0b01vZGVsJTBBJTBBdHJhbnNmb3JtZXIlMjAlM0QlMjBBdXRvTW9kZWwuZnJvbV9wcmV0cmFpbmVkKCUwQSUyMCUyMCUyMCUyMCUyMk92ZXJ3b3JsZCUyRldheXBvaW50LTEtU21hbGwlMjIlMkMlMjBzdWJmb2xkZXIlM0QlMjJ0cmFuc2Zvcm1lciUyMiUyQyUyMHRydXN0X3JlbW90ZV9jb2RlJTNEVHJ1ZSUyQyUyMGR0eXBlJTNEdG9yY2guYmZsb2F0MTYlMkMlMjBkZXZpY2VfbWFwJTNEJTIyY3VkYSUyMiUwQSk=",highlighted:`<span class="hljs-keyword">import</span> torch
<span class="hljs-keyword">from</span> diffusers <span class="hljs-keyword">import</span> AutoModel
transformer = AutoModel.from_pretrained(
<span class="hljs-string">&quot;Overworld/Waypoint-1-Small&quot;</span>, subfolder=<span class="hljs-string">&quot;transformer&quot;</span>, trust_remote_code=<span class="hljs-literal">True</span>, dtype=torch.bfloat16, device_map=<span class="hljs-string">&quot;cuda&quot;</span>
)`,lang:"py",wrap:!1});var o=s(u,4),b=s(N(o),2);l(b,{code:"dHJhbnNmb3JtZXIlMjAlM0QlMjBBdXRvTW9kZWwuZnJvbV9wcmV0cmFpbmVkKCUwQSUyMCUyMCUyMCUyMCUyMk92ZXJ3b3JsZCUyRldheXBvaW50LTEtU21hbGwlMjIlMkMlMjBzdWJmb2xkZXIlM0QlMjJ0cmFuc2Zvcm1lciUyMiUyQyUyMHRydXN0X3JlbW90ZV9jb2RlJTNEVHJ1ZSUyQyUyMHJldmlzaW9uJTNEJTIyYTNkOGNiMiUyMiUwQSk=",highlighted:`transformer = AutoModel.from_pretrained(
<span class="hljs-string">&quot;Overworld/Waypoint-1-Small&quot;</span>, subfolder=<span class="hljs-string">&quot;transformer&quot;</span>, trust_remote_code=<span class="hljs-literal">True</span>, revision=<span class="hljs-string">&quot;a3d8cb2&quot;</span>
)`,lang:"py",wrap:!1}),X(o);var M=s(o,2);a(M,{title:"Saving custom models",local:"saving-custom-models",headingTag:"h3"});var y=s(M,4);l(y,{code:"JTIzJTIwbXlfbW9kZWwucHklMEFmcm9tJTIwZGlmZnVzZXJzJTIwaW1wb3J0JTIwTW9kZWxNaXhpbiUyQyUyMENvbmZpZ01peGluJTBBJTBBY2xhc3MlMjBNeUN1c3RvbU1vZGVsKE1vZGVsTWl4aW4lMkMlMjBDb25maWdNaXhpbiklM0ElMEElMjAlMjAlMjAlMjAuLi4lMEElMEFNeUN1c3RvbU1vZGVsLnJlZ2lzdGVyX2Zvcl9hdXRvX2NsYXNzKCUyMkF1dG9Nb2RlbCUyMiklMEElMEFtb2RlbCUyMCUzRCUyME15Q3VzdG9tTW9kZWwoLi4uKSUwQW1vZGVsLnNhdmVfcHJldHJhaW5lZCglMjIuJTJGbXlfbW9kZWwlMjIp",highlighted:`<span class="hljs-comment"># my_model.py</span>
<span class="hljs-keyword">from</span> diffusers <span class="hljs-keyword">import</span> ModelMixin, ConfigMixin
<span class="hljs-keyword">class</span> <span class="hljs-title class_">MyCustomModel</span>(ModelMixin, ConfigMixin):
...
MyCustomModel.register_for_auto_class(<span class="hljs-string">&quot;AutoModel&quot;</span>)
model = MyCustomModel(...)
model.save_pretrained(<span class="hljs-string">&quot;./my_model&quot;</span>)`,lang:"py",wrap:!1});var h=s(y,4);l(h,{code:"JTdCJTBBJTIwJTIwJTIyYXV0b19tYXAlMjIlM0ElMjAlN0IlMEElMjAlMjAlMjAlMjAlMjJBdXRvTW9kZWwlMjIlM0ElMjAlMjJteV9tb2RlbC5NeUN1c3RvbU1vZGVsJTIyJTBBJTIwJTIwJTdEJTBBJTdE",highlighted:`<span class="hljs-punctuation">{</span>
<span class="hljs-attr">&quot;auto_map&quot;</span><span class="hljs-punctuation">:</span> <span class="hljs-punctuation">{</span>
<span class="hljs-attr">&quot;AutoModel&quot;</span><span class="hljs-punctuation">:</span> <span class="hljs-string">&quot;my_model.MyCustomModel&quot;</span>
<span class="hljs-punctuation">}</span>
<span class="hljs-punctuation">}</span>`,lang:"json",wrap:!1});var g=s(h,4);_(g,{source:"https://github.com/huggingface/diffusers/blob/main/docs/source/en/using-diffusers/automodel.md"}),R(2),T(w,e),B()}export{Y as component};

Xet Storage Details

Size:
11.8 kB
·
Xet hash:
6baef7a05b0bfab5382119176cec6e81d45f86ede83046cd8f0375c96c15ee6a

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