Buckets:

HuggingFaceDocBuilder's picture
download
raw
39.1 kB
import"../chunks/DsnmJJEf.js";import{i as S,h as K,C as $,H as p,a as j,D as e,b as a,E as ss,s as as}from"../chunks/Cd6DUeqq.js";import{p as ts,o as ns,s,f as L,a as y,b as es,d as t,c as b,n as m,r as n}from"../chunks/BLjEbp4u.js";import{E as ls}from"../chunks/OTCdYc08.js";const ps='{"title":"BEMA for Reference Model","local":"bema-for-reference-model","sections":[{"title":"Usage","local":"usage","sections":[],"depth":2},{"title":"DPOTrainer","local":"trl.DPOTrainer","sections":[],"depth":2},{"title":"BEMACallback","local":"trl.BEMACallback","sections":[],"depth":2}],"depth":1}';var ms=b('<meta name="hf:doc:metadata"/>'),is=b("<p>Example:</p> <!>",1),os=b(`<p></p> <!> <!> <p>This feature implements the BEMA algorithm to update the reference model during DPO training.</p> <!> <!> <!> <div class="docstring border-l-2 border-t-2 pl-4 pt-3.5 border-gray-100 rounded-tl-xl mb-6 mt-8"><!> <div class="docstring border-l-2 border-t-2 pl-4 pt-3.5 border-gray-100 rounded-tl-xl mb-6 mt-8"><!> <p>Main training entry point.</p></div> <div class="docstring border-l-2 border-t-2 pl-4 pt-3.5 border-gray-100 rounded-tl-xl mb-6 mt-8"><!> <p>Will save the model, so you can reload it using <code>from_pretrained()</code>.</p> <p>Will only save from the main process.</p></div> <div class="docstring border-l-2 border-t-2 pl-4 pt-3.5 border-gray-100 rounded-tl-xl mb-6 mt-8"><!> <p>Upload <code>self.model</code> and <code>self.processing_class</code> to the 🤗 model hub on the repo <code>self.args.hub_model_id</code>.</p></div></div> <!> <div class="docstring border-l-2 border-t-2 pl-4 pt-3.5 border-gray-100 rounded-tl-xl mb-6 mt-8"><!> <p>A <a href="https://huggingface.co/docs/transformers/main/en/main_classes/callback#transformers.TrainerCallback" rel="nofollow">TrainerCallback</a> that implements <a href="https://huggingface.co/papers/2508.00180" rel="nofollow">BEMA</a> (Bias-Corrected Exponential Moving Average) by <a href="https://huggingface.co/abblock" rel="nofollow">Adam Block</a> and <a href="https://huggingface.co/cyrilzhang" rel="nofollow">Cyril
Zhang</a>. Code from <a href="https://github.com/abblock/bema" rel="nofollow">https://github.com/abblock/bema</a> under MIT license.</p> <p>BEMA computes model weights that scale like: <!></p> <p>where <!> is the current model weights, <!> is a snapshot of the model weights at the
first <code>update_after</code> step, <!> is the exponential moving average of the model weights, and<!> is a scaling factor that decays with the number of steps <!> as <!></p> <p>The EMA is computed as: <!></p> <p>where <!> is a decay factor that decays with the number of steps <!> as <!></p> <!></div> <!> <p></p>`,1);function ds(W,Z){ts(Z,!1),ns(()=>{new URLSearchParams(window.location.search).get("fw")}),S();var f=os();K("1hwg6o0",l=>{var v=ms();as(v,"content",ps),y(l,v)});var w=s(L(f),2);$(w,{containerStyle:"float: right; margin-left: 10px; display: inline-flex; position: relative; z-index: 10;"});var x=s(w,2);p(x,{title:"BEMA for Reference Model",local:"bema-for-reference-model",headingTag:"h1"});var _=s(x,4);p(_,{title:"Usage",local:"usage",headingTag:"h2"});var M=s(_,2);j(M,{code:"ZnJvbSUyMHRybC5leHBlcmltZW50YWwuYmVtYV9mb3JfcmVmX21vZGVsJTIwaW1wb3J0JTIwQkVNQUNhbGxiYWNrJTJDJTIwRFBPVHJhaW5lciUwQWZyb20lMjBkYXRhc2V0cyUyMGltcG9ydCUyMGxvYWRfZGF0YXNldCUwQSUwQWRhdGFzZXQlMjAlM0QlMjBsb2FkX2RhdGFzZXQoJTIydHJsLWludGVybmFsLXRlc3RpbmclMkZ6ZW4lMjIlMkMlMjAlMjJzdGFuZGFyZF9wcmVmZXJlbmNlJTIyJTJDJTIwc3BsaXQlM0QlMjJ0cmFpbiUyMiklMEElMEFiZW1hX2NhbGxiYWNrJTIwJTNEJTIwQkVNQUNhbGxiYWNrKHVwZGF0ZV9yZWZfbW9kZWwlM0RUcnVlKSUwQSUwQXRyYWluZXIlMjAlM0QlMjBEUE9UcmFpbmVyKCUwQSUyMCUyMCUyMCUyMG1vZGVsJTNEJTIydHJsLWludGVybmFsLXRlc3RpbmclMkZ0aW55LVF3ZW4yRm9yQ2F1c2FsTE0tMi41JTIyJTJDJTBBJTIwJTIwJTIwJTIwdHJhaW5fZGF0YXNldCUzRGRhdGFzZXQlMkMlMEElMjAlMjAlMjAlMjBjYWxsYmFja3MlM0QlNUJiZW1hX2NhbGxiYWNrJTVEJTJDJTBBKSUwQXRyYWluZXIudHJhaW4oKQ==",highlighted:`<span class="hljs-keyword">from</span> trl.experimental.bema_for_ref_model <span class="hljs-keyword">import</span> BEMACallback, DPOTrainer
<span class="hljs-keyword">from</span> datasets <span class="hljs-keyword">import</span> load_dataset
dataset = load_dataset(<span class="hljs-string">&quot;trl-internal-testing/zen&quot;</span>, <span class="hljs-string">&quot;standard_preference&quot;</span>, split=<span class="hljs-string">&quot;train&quot;</span>)
bema_callback = BEMACallback(update_ref_model=<span class="hljs-literal">True</span>)
trainer = DPOTrainer(
model=<span class="hljs-string">&quot;trl-internal-testing/tiny-Qwen2ForCausalLM-2.5&quot;</span>,
train_dataset=dataset,
callbacks=[bema_callback],
)
trainer.train()`,lang:"python",wrap:!1});var q=s(M,2);p(q,{title:"DPOTrainer",local:"trl.DPOTrainer",headingTag:"h2"});var i=s(q,2),k=t(i);e(k,{name:"class trl.DPOTrainer",anchor:"trl.DPOTrainer",source:"https://github.com/huggingface/trl/blob/vr_5313/trl/experimental/bema_for_ref_model/dpo_trainer.py#L19",parameters:[{name:"*args",val:""},{name:"**kwargs",val:""}]});var o=s(k,2),Q=t(o);e(Q,{name:"train",anchor:"trl.DPOTrainer.train",source:"https://github.com/huggingface/trl/blob/vr_5313/transformers/trainer.py#L1331",parameters:[{name:"resume_from_checkpoint",val:": str | bool | None = None"},{name:"trial",val:": optuna.Trial | dict[str, Any] | None = None"},{name:"ignore_keys_for_eval",val:": list[str] | None = None"}],parametersDescription:[{anchor:"trl.DPOTrainer.train.resume_from_checkpoint",description:`<strong>resume_from_checkpoint</strong> (<code>str</code> or <code>bool</code>, <em>optional</em>) &#x2014;
If a <code>str</code>, local path to a saved checkpoint as saved by a previous instance of <code>Trainer</code>. If a
<code>bool</code> and equals <code>True</code>, load the last checkpoint in <em>args.output_dir</em> as saved by a previous instance
of <code>Trainer</code>. If present, training will resume from the model/optimizer/scheduler states loaded here.`,name:"resume_from_checkpoint"},{anchor:"trl.DPOTrainer.train.trial",description:`<strong>trial</strong> (<code>optuna.Trial</code> or <code>dict[str, Any]</code>, <em>optional</em>) &#x2014;
The trial run or the hyperparameter dictionary for hyperparameter search.`,name:"trial"},{anchor:"trl.DPOTrainer.train.ignore_keys_for_eval",description:`<strong>ignore_keys_for_eval</strong> (<code>list[str]</code>, <em>optional</em>) &#x2014;
A list of keys in the output of your model (if it is a dictionary) that should be ignored when
gathering predictions for evaluation during the training.`,name:"ignore_keys_for_eval"}],returnDescription:`<script context="module">export const metadata = 'undefined';<\/script>
<p>Object containing the global step count, training loss, and metrics.</p>
`,returnType:`<script context="module">export const metadata = 'undefined';<\/script>
<p><code>~trainer_utils.TrainOutput</code></p>
`}),m(2),n(o);var r=s(o,2),P=t(r);e(P,{name:"save_model",anchor:"trl.DPOTrainer.save_model",source:"https://github.com/huggingface/trl/blob/vr_5313/transformers/trainer.py#L3775",parameters:[{name:"output_dir",val:": str | None = None"},{name:"_internal_call",val:": bool = False"}]}),m(4),n(r);var T=s(r,2),G=t(T);e(G,{name:"push_to_hub",anchor:"trl.DPOTrainer.push_to_hub",source:"https://github.com/huggingface/trl/blob/vr_5313/transformers/trainer.py#L4022",parameters:[{name:"commit_message",val:": str | None = 'End of training'"},{name:"blocking",val:": bool = True"},{name:"token",val:": str | None = None"},{name:"revision",val:": str | None = None"},{name:"**kwargs",val:""}],parametersDescription:[{anchor:"trl.DPOTrainer.push_to_hub.commit_message",description:`<strong>commit_message</strong> (<code>str</code>, <em>optional</em>, defaults to <code>&quot;End of training&quot;</code>) &#x2014;
Message to commit while pushing.`,name:"commit_message"},{anchor:"trl.DPOTrainer.push_to_hub.blocking",description:`<strong>blocking</strong> (<code>bool</code>, <em>optional</em>, defaults to <code>True</code>) &#x2014;
Whether the function should return only when the <code>git push</code> has finished.`,name:"blocking"},{anchor:"trl.DPOTrainer.push_to_hub.token",description:`<strong>token</strong> (<code>str</code>, <em>optional</em>, defaults to <code>None</code>) &#x2014;
Token with write permission to overwrite Trainer&#x2019;s original args.`,name:"token"},{anchor:"trl.DPOTrainer.push_to_hub.revision",description:`<strong>revision</strong> (<code>str</code>, <em>optional</em>) &#x2014;
The git revision to commit from. Defaults to the head of the &#x201C;main&#x201D; branch.`,name:"revision"},{anchor:"trl.DPOTrainer.push_to_hub.kwargs",description:`<strong>kwargs</strong> (<code>dict[str, Any]</code>, <em>optional</em>) &#x2014;
Additional keyword arguments passed along to <code>~Trainer.create_model_card</code>.`,name:"kwargs"}],returnDescription:`<script context="module">export const metadata = 'undefined';<\/script>
<p>The URL of the repository where the model was pushed if <code>blocking=False</code>, or a <code>Future</code> object tracking the
progress of the commit if <code>blocking=True</code>.</p>
`}),m(2),n(T),n(i);var E=s(i,2);p(E,{title:"BEMACallback",local:"trl.BEMACallback",headingTag:"h2"});var c=s(E,2),z=t(c);e(z,{name:"class trl.BEMACallback",anchor:"trl.BEMACallback",source:"https://github.com/huggingface/trl/blob/vr_5313/trl/experimental/bema_for_ref_model/callback.py#L59",parameters:[{name:"update_freq",val:": int = 400"},{name:"ema_power",val:": float = 0.5"},{name:"bias_power",val:": float = 0.2"},{name:"lag",val:": int = 10"},{name:"update_after",val:": int = 0"},{name:"multiplier",val:": float = 1.0"},{name:"min_ema_multiplier",val:": float = 0.0"},{name:"device",val:": str = 'cpu'"},{name:"update_ref_model",val:": bool = False"},{name:"ref_model_update_freq",val:": int = 400"},{name:"ref_model_update_after",val:": int = 0"}],parametersDescription:[{anchor:"trl.BEMACallback.update_freq",description:`<strong>update_freq</strong> (<code>int</code>, <em>optional</em>, defaults to <code>400</code>) &#x2014;
Update the BEMA weights every X steps. Denoted this as {@html &quot;<span class="\\&quot;katex\\&quot;"><span class="\\&quot;katex-mathml\\&quot;"><math xmlns="\\&quot;http://www.w3.org/1998/Math/MathML\\&quot;"><semantics><mrow><mi>&#x3D5;</mi></mrow><annotation encoding="\\&quot;application/x-tex\\&quot;"> \\\\phi </annotation></semantics></math></span><span class="\\&quot;katex-html\\&quot;" aria-hidden="\\&quot;true\\&quot;"><span class="\\&quot;base\\&quot;"><span class="\\&quot;strut\\&quot;" style="\\&quot;height:0.8889em;vertical-align:-0.1944em;\\&quot;"></span><span class="\\&quot;mord" mathnormal\\">&#x3D5;</span></span></span></span>&quot;} in the paper.`,name:"update_freq"},{anchor:"trl.BEMACallback.ema_power",description:`<strong>ema_power</strong> (<code>float</code>, <em>optional</em>, defaults to <code>0.5</code>) &#x2014;
Power for the EMA decay factor. Denoted {@html &quot;<span class="\\&quot;katex\\&quot;"><span class="\\&quot;katex-mathml\\&quot;"><math xmlns="\\&quot;http://www.w3.org/1998/Math/MathML\\&quot;"><semantics><mrow><mi>&#x3BA;</mi></mrow><annotation encoding="\\&quot;application/x-tex\\&quot;"> \\\\kappa </annotation></semantics></math></span><span class="\\&quot;katex-html\\&quot;" aria-hidden="\\&quot;true\\&quot;"><span class="\\&quot;base\\&quot;"><span class="\\&quot;strut\\&quot;" style="\\&quot;height:0.4306em;\\&quot;"></span><span class="\\&quot;mord" mathnormal\\">&#x3BA;</span></span></span></span>&quot;} in the paper. To disable EMA, set this to <code>0.0</code>.`,name:"ema_power"},{anchor:"trl.BEMACallback.bias_power",description:`<strong>bias_power</strong> (<code>float</code>, <em>optional</em>, defaults to <code>0.2</code>) &#x2014;
Power for the BEMA scaling factor. Denoted {@html &quot;<span class="\\&quot;katex\\&quot;"><span class="\\&quot;katex-mathml\\&quot;"><math xmlns="\\&quot;http://www.w3.org/1998/Math/MathML\\&quot;"><semantics><mrow><mi>&#x3B7;</mi></mrow><annotation encoding="\\&quot;application/x-tex\\&quot;"> \\\\eta </annotation></semantics></math></span><span class="\\&quot;katex-html\\&quot;" aria-hidden="\\&quot;true\\&quot;"><span class="\\&quot;base\\&quot;"><span class="\\&quot;strut\\&quot;" style="\\&quot;height:0.625em;vertical-align:-0.1944em;\\&quot;"></span><span class="\\&quot;mord" mathnormal\\" style="\\&quot;margin-right:0.0359em;\\&quot;">&#x3B7;</span></span></span></span>&quot;} in the paper. To disable BEMA, set this to <code>0.0</code>.`,name:"bias_power"},{anchor:"trl.BEMACallback.lag",description:`<strong>lag</strong> (<code>int</code>, <em>optional</em>, defaults to <code>10</code>) &#x2014;
Initial offset in the weight decay schedule that controls early-stage smoothness by acting as a virtual
starting age for the updates. Denoted as {@html &quot;<span class="\\&quot;katex\\&quot;"><span class="\\&quot;katex-mathml\\&quot;"><math xmlns="\\&quot;http://www.w3.org/1998/Math/MathML\\&quot;"><semantics><mrow><mi>&#x3C1;</mi></mrow><annotation encoding="\\&quot;application/x-tex\\&quot;"> \\\\rho </annotation></semantics></math></span><span class="\\&quot;katex-html\\&quot;" aria-hidden="\\&quot;true\\&quot;"><span class="\\&quot;base\\&quot;"><span class="\\&quot;strut\\&quot;" style="\\&quot;height:0.625em;vertical-align:-0.1944em;\\&quot;"></span><span class="\\&quot;mord" mathnormal\\">&#x3C1;</span></span></span></span>&quot;} in the paper.`,name:"lag"},{anchor:"trl.BEMACallback.update_after",description:`<strong>update_after</strong> (<code>int</code>, <em>optional</em>, defaults to <code>0</code>) &#x2014;
Burn-in time before starting to update the BEMA weights. Denoted {@html &quot;<span class="\\&quot;katex\\&quot;"><span class="\\&quot;katex-mathml\\&quot;"><math xmlns="\\&quot;http://www.w3.org/1998/Math/MathML\\&quot;"><semantics><mrow><mi>&#x3C4;</mi></mrow><annotation encoding="\\&quot;application/x-tex\\&quot;"> \\\\tau </annotation></semantics></math></span><span class="\\&quot;katex-html\\&quot;" aria-hidden="\\&quot;true\\&quot;"><span class="\\&quot;base\\&quot;"><span class="\\&quot;strut\\&quot;" style="\\&quot;height:0.4306em;\\&quot;"></span><span class="\\&quot;mord" mathnormal\\" style="\\&quot;margin-right:0.1132em;\\&quot;">&#x3C4;</span></span></span></span>&quot;} in the paper.`,name:"update_after"},{anchor:"trl.BEMACallback.multiplier",description:`<strong>multiplier</strong> (<code>float</code>, <em>optional</em>, defaults to <code>1.0</code>) &#x2014;
Initial value for the EMA decay factor. Denoted as {@html &quot;<span class="\\&quot;katex\\&quot;"><span class="\\&quot;katex-mathml\\&quot;"><math xmlns="\\&quot;http://www.w3.org/1998/Math/MathML\\&quot;"><semantics><mrow><mi>&#x3B3;</mi></mrow><annotation encoding="\\&quot;application/x-tex\\&quot;"> \\\\gamma </annotation></semantics></math></span><span class="\\&quot;katex-html\\&quot;" aria-hidden="\\&quot;true\\&quot;"><span class="\\&quot;base\\&quot;"><span class="\\&quot;strut\\&quot;" style="\\&quot;height:0.625em;vertical-align:-0.1944em;\\&quot;"></span><span class="\\&quot;mord" mathnormal\\" style="\\&quot;margin-right:0.0556em;\\&quot;">&#x3B3;</span></span></span></span>&quot;} in the paper.`,name:"multiplier"},{anchor:"trl.BEMACallback.min_ema_multiplier",description:`<strong>min_ema_multiplier</strong> (<code>float</code>, <em>optional</em>, defaults to <code>0.0</code>) &#x2014;
Minimum value for the EMA decay factor.`,name:"min_ema_multiplier"},{anchor:"trl.BEMACallback.device",description:`<strong>device</strong> (<code>str</code>, <em>optional</em>, defaults to <code>&quot;cpu&quot;</code>) &#x2014;
Device to use for the BEMA buffers, e.g. <code>&quot;cpu&quot;</code> or <code>&quot;cuda&quot;</code>. Note that in most cases, this device SHOULD
BE DIFFERENT from the device used for training in order to avoid OOM.`,name:"device"},{anchor:"trl.BEMACallback.update_ref_model",description:`<strong>update_ref_model</strong> (<code>bool</code>, <em>optional</em>, defaults to <code>False</code>) &#x2014;
Whether to update the reference model with BEMA weights. This creates a lagged, smoothed version of the
main model as the reference model.`,name:"update_ref_model"},{anchor:"trl.BEMACallback.ref_model_update_freq",description:`<strong>ref_model_update_freq</strong> (<code>int</code>, <em>optional</em>, defaults to <code>400</code>) &#x2014;
Update the reference model with BEMA weights every this many steps.`,name:"ref_model_update_freq"},{anchor:"trl.BEMACallback.ref_model_update_after",description:`<strong>ref_model_update_after</strong> (<code>int</code>, <em>optional</em>, defaults to <code>0</code>) &#x2014;
Number of steps to wait before starting to update the reference model.`,name:"ref_model_update_after"}]});var h=s(z,4),O=s(t(h));a(O,()=>`<span class="katex-display"><span class="katex"><span class="katex-mathml"><math xmlns="http://www.w3.org/1998/Math/MathML" display="block"><semantics><mrow><msubsup><mi>θ</mi><mi>t</mi><mo mathvariant="normal" lspace="0em" rspace="0em">′</mo></msubsup><mo>=</mo><msub><mi>α</mi><mi>t</mi></msub><mo>⋅</mo><mo stretchy="false">(</mo><msub><mi>θ</mi><mi>t</mi></msub><mo>−</mo><msub><mi>θ</mi><mn>0</mn></msub><mo stretchy="false">)</mo><mo>+</mo><msub><mtext>EMA</mtext><mi>t</mi></msub></mrow><annotation encoding="application/x-tex">
\\theta_t&#x27; = \\alpha_t \\cdot (\\theta_t - \\theta_0) + \\text{EMA}_t
</annotation></semantics></math></span><span class="katex-html" aria-hidden="true"><span class="base"><span class="strut" style="height:1.0489em;vertical-align:-0.247em;"></span><span class="mord"><span class="mord mathnormal" style="margin-right:0.0278em;">θ</span><span class="msupsub"><span class="vlist-t vlist-t2"><span class="vlist-r"><span class="vlist" style="height:0.8019em;"><span style="top:-2.453em;margin-left:-0.0278em;margin-right:0.05em;"><span class="pstrut" style="height:2.7em;"></span><span class="sizing reset-size6 size3 mtight"><span class="mord mathnormal mtight">t</span></span></span><span style="top:-3.113em;margin-right:0.05em;"><span class="pstrut" style="height:2.7em;"></span><span class="sizing reset-size6 size3 mtight"><span class="mord mtight"><span class="mord mtight">′</span></span></span></span></span><span class="vlist-s">​</span></span><span class="vlist-r"><span class="vlist" style="height:0.247em;"><span></span></span></span></span></span></span><span class="mspace" style="margin-right:0.2778em;"></span><span class="mrel">=</span><span class="mspace" style="margin-right:0.2778em;"></span></span><span class="base"><span class="strut" style="height:0.5945em;vertical-align:-0.15em;"></span><span class="mord"><span class="mord mathnormal" style="margin-right:0.0037em;">α</span><span class="msupsub"><span class="vlist-t vlist-t2"><span class="vlist-r"><span class="vlist" style="height:0.2806em;"><span style="top:-2.55em;margin-left:-0.0037em;margin-right:0.05em;"><span class="pstrut" style="height:2.7em;"></span><span class="sizing reset-size6 size3 mtight"><span class="mord mathnormal mtight">t</span></span></span></span><span class="vlist-s">​</span></span><span class="vlist-r"><span class="vlist" style="height:0.15em;"><span></span></span></span></span></span></span><span class="mspace" style="margin-right:0.2222em;"></span><span class="mbin">⋅</span><span class="mspace" style="margin-right:0.2222em;"></span></span><span class="base"><span class="strut" style="height:1em;vertical-align:-0.25em;"></span><span class="mopen">(</span><span class="mord"><span class="mord mathnormal" style="margin-right:0.0278em;">θ</span><span class="msupsub"><span class="vlist-t vlist-t2"><span class="vlist-r"><span class="vlist" style="height:0.2806em;"><span style="top:-2.55em;margin-left:-0.0278em;margin-right:0.05em;"><span class="pstrut" style="height:2.7em;"></span><span class="sizing reset-size6 size3 mtight"><span class="mord mathnormal mtight">t</span></span></span></span><span class="vlist-s">​</span></span><span class="vlist-r"><span class="vlist" style="height:0.15em;"><span></span></span></span></span></span></span><span class="mspace" style="margin-right:0.2222em;"></span><span class="mbin">−</span><span class="mspace" style="margin-right:0.2222em;"></span></span><span class="base"><span class="strut" style="height:1em;vertical-align:-0.25em;"></span><span class="mord"><span class="mord mathnormal" style="margin-right:0.0278em;">θ</span><span class="msupsub"><span class="vlist-t vlist-t2"><span class="vlist-r"><span class="vlist" style="height:0.3011em;"><span style="top:-2.55em;margin-left:-0.0278em;margin-right:0.05em;"><span class="pstrut" style="height:2.7em;"></span><span class="sizing reset-size6 size3 mtight"><span class="mord mtight">0</span></span></span></span><span class="vlist-s">​</span></span><span class="vlist-r"><span class="vlist" style="height:0.15em;"><span></span></span></span></span></span></span><span class="mclose">)</span><span class="mspace" style="margin-right:0.2222em;"></span><span class="mbin">+</span><span class="mspace" style="margin-right:0.2222em;"></span></span><span class="base"><span class="strut" style="height:0.8333em;vertical-align:-0.15em;"></span><span class="mord"><span class="mord text"><span class="mord">EMA</span></span><span class="msupsub"><span class="vlist-t vlist-t2"><span class="vlist-r"><span class="vlist" style="height:0.2806em;"><span style="top:-2.55em;margin-right:0.05em;"><span class="pstrut" style="height:2.7em;"></span><span class="sizing reset-size6 size3 mtight"><span class="mord mathnormal mtight">t</span></span></span></span><span class="vlist-s">​</span></span><span class="vlist-r"><span class="vlist" style="height:0.15em;"><span></span></span></span></span></span></span></span></span></span></span>`),n(h);var g=s(h,2),A=s(t(g));a(A,()=>'<span class="katex"><span class="katex-mathml"><math xmlns="http://www.w3.org/1998/Math/MathML"><semantics><mrow><msub><mi>θ</mi><mi>t</mi></msub></mrow><annotation encoding="application/x-tex"> \\theta_t </annotation></semantics></math></span><span class="katex-html" aria-hidden="true"><span class="base"><span class="strut" style="height:0.8444em;vertical-align:-0.15em;"></span><span class="mord"><span class="mord mathnormal" style="margin-right:0.0278em;">θ</span><span class="msupsub"><span class="vlist-t vlist-t2"><span class="vlist-r"><span class="vlist" style="height:0.2806em;"><span style="top:-2.55em;margin-left:-0.0278em;margin-right:0.05em;"><span class="pstrut" style="height:2.7em;"></span><span class="sizing reset-size6 size3 mtight"><span class="mord mathnormal mtight">t</span></span></span></span><span class="vlist-s">​</span></span><span class="vlist-r"><span class="vlist" style="height:0.15em;"><span></span></span></span></span></span></span></span></span></span>');var B=s(A,2);a(B,()=>'<span class="katex"><span class="katex-mathml"><math xmlns="http://www.w3.org/1998/Math/MathML"><semantics><mrow><msub><mi>θ</mi><mn>0</mn></msub></mrow><annotation encoding="application/x-tex"> \\theta_0 </annotation></semantics></math></span><span class="katex-html" aria-hidden="true"><span class="base"><span class="strut" style="height:0.8444em;vertical-align:-0.15em;"></span><span class="mord"><span class="mord mathnormal" style="margin-right:0.0278em;">θ</span><span class="msupsub"><span class="vlist-t vlist-t2"><span class="vlist-r"><span class="vlist" style="height:0.3011em;"><span style="top:-2.55em;margin-left:-0.0278em;margin-right:0.05em;"><span class="pstrut" style="height:2.7em;"></span><span class="sizing reset-size6 size3 mtight"><span class="mord mtight">0</span></span></span></span><span class="vlist-s">​</span></span><span class="vlist-r"><span class="vlist" style="height:0.15em;"><span></span></span></span></span></span></span></span></span></span>');var C=s(B,4);a(C,()=>'<span class="katex"><span class="katex-mathml"><math xmlns="http://www.w3.org/1998/Math/MathML"><semantics><mrow><msub><mtext>EMA</mtext><mi>t</mi></msub></mrow><annotation encoding="application/x-tex"> \\text{EMA}_t </annotation></semantics></math></span><span class="katex-html" aria-hidden="true"><span class="base"><span class="strut" style="height:0.8333em;vertical-align:-0.15em;"></span><span class="mord"><span class="mord text"><span class="mord">EMA</span></span><span class="msupsub"><span class="vlist-t vlist-t2"><span class="vlist-r"><span class="vlist" style="height:0.2806em;"><span style="top:-2.55em;margin-right:0.05em;"><span class="pstrut" style="height:2.7em;"></span><span class="sizing reset-size6 size3 mtight"><span class="mord mathnormal mtight">t</span></span></span></span><span class="vlist-s">​</span></span><span class="vlist-r"><span class="vlist" style="height:0.15em;"><span></span></span></span></span></span></span></span></span></span>');var D=s(C,2);a(D,()=>'<span class="katex"><span class="katex-mathml"><math xmlns="http://www.w3.org/1998/Math/MathML"><semantics><mrow><msub><mi>α</mi><mi>t</mi></msub></mrow><annotation encoding="application/x-tex"> \\alpha_t </annotation></semantics></math></span><span class="katex-html" aria-hidden="true"><span class="base"><span class="strut" style="height:0.5806em;vertical-align:-0.15em;"></span><span class="mord"><span class="mord mathnormal" style="margin-right:0.0037em;">α</span><span class="msupsub"><span class="vlist-t vlist-t2"><span class="vlist-r"><span class="vlist" style="height:0.2806em;"><span style="top:-2.55em;margin-left:-0.0037em;margin-right:0.05em;"><span class="pstrut" style="height:2.7em;"></span><span class="sizing reset-size6 size3 mtight"><span class="mord mathnormal mtight">t</span></span></span></span><span class="vlist-s">​</span></span><span class="vlist-r"><span class="vlist" style="height:0.15em;"><span></span></span></span></span></span></span></span></span></span>');var J=s(D,2);a(J,()=>'<span class="katex"><span class="katex-mathml"><math xmlns="http://www.w3.org/1998/Math/MathML"><semantics><mrow><mi>t</mi></mrow><annotation encoding="application/x-tex"> t </annotation></semantics></math></span><span class="katex-html" aria-hidden="true"><span class="base"><span class="strut" style="height:0.6151em;"></span><span class="mord mathnormal">t</span></span></span></span>');var I=s(J,2);a(I,()=>`<span class="katex-display"><span class="katex"><span class="katex-mathml"><math xmlns="http://www.w3.org/1998/Math/MathML" display="block"><semantics><mrow><msub><mi>α</mi><mi>t</mi></msub><mo>=</mo><mo stretchy="false">(</mo><mi>ρ</mi><mo>+</mo><mi>γ</mi><mo>⋅</mo><mi>t</mi><msup><mo stretchy="false">)</mo><mrow><mo>−</mo><mi>η</mi></mrow></msup><mi mathvariant="normal">.</mi></mrow><annotation encoding="application/x-tex">
\\alpha_t = (\\rho + \\gamma \\cdot t)^{-\\eta}.
</annotation></semantics></math></span><span class="katex-html" aria-hidden="true"><span class="base"><span class="strut" style="height:0.5806em;vertical-align:-0.15em;"></span><span class="mord"><span class="mord mathnormal" style="margin-right:0.0037em;">α</span><span class="msupsub"><span class="vlist-t vlist-t2"><span class="vlist-r"><span class="vlist" style="height:0.2806em;"><span style="top:-2.55em;margin-left:-0.0037em;margin-right:0.05em;"><span class="pstrut" style="height:2.7em;"></span><span class="sizing reset-size6 size3 mtight"><span class="mord mathnormal mtight">t</span></span></span></span><span class="vlist-s">​</span></span><span class="vlist-r"><span class="vlist" style="height:0.15em;"><span></span></span></span></span></span></span><span class="mspace" style="margin-right:0.2778em;"></span><span class="mrel">=</span><span class="mspace" style="margin-right:0.2778em;"></span></span><span class="base"><span class="strut" style="height:1em;vertical-align:-0.25em;"></span><span class="mopen">(</span><span class="mord mathnormal">ρ</span><span class="mspace" style="margin-right:0.2222em;"></span><span class="mbin">+</span><span class="mspace" style="margin-right:0.2222em;"></span></span><span class="base"><span class="strut" style="height:0.6389em;vertical-align:-0.1944em;"></span><span class="mord mathnormal" style="margin-right:0.0556em;">γ</span><span class="mspace" style="margin-right:0.2222em;"></span><span class="mbin">⋅</span><span class="mspace" style="margin-right:0.2222em;"></span></span><span class="base"><span class="strut" style="height:1.0713em;vertical-align:-0.25em;"></span><span class="mord mathnormal">t</span><span class="mclose"><span class="mclose">)</span><span class="msupsub"><span class="vlist-t"><span class="vlist-r"><span class="vlist" style="height:0.8213em;"><span style="top:-3.113em;margin-right:0.05em;"><span class="pstrut" style="height:2.7em;"></span><span class="sizing reset-size6 size3 mtight"><span class="mord mtight"><span class="mord mtight">−</span><span class="mord mathnormal mtight" style="margin-right:0.0359em;">η</span></span></span></span></span></span></span></span></span><span class="mord">.</span></span></span></span></span>`),n(g);var d=s(g,2),R=s(t(d));a(R,()=>`<span class="katex-display"><span class="katex"><span class="katex-mathml"><math xmlns="http://www.w3.org/1998/Math/MathML" display="block"><semantics><mrow><msub><mtext>EMA</mtext><mi>t</mi></msub><mo>=</mo><mo stretchy="false">(</mo><mn>1</mn><mo>−</mo><msub><mi>β</mi><mi>t</mi></msub><mo stretchy="false">)</mo><mo>⋅</mo><msub><mtext>EMA</mtext><mrow><mi>t</mi><mo>−</mo><mn>1</mn></mrow></msub><mo>+</mo><msub><mi>β</mi><mi>t</mi></msub><mo>⋅</mo><msub><mi>θ</mi><mi>t</mi></msub></mrow><annotation encoding="application/x-tex">
\\text{EMA}_t = (1 - \\beta_t) \\cdot \\text{EMA}_{t-1} + \\beta_t \\cdot \\theta_t
</annotation></semantics></math></span><span class="katex-html" aria-hidden="true"><span class="base"><span class="strut" style="height:0.8333em;vertical-align:-0.15em;"></span><span class="mord"><span class="mord text"><span class="mord">EMA</span></span><span class="msupsub"><span class="vlist-t vlist-t2"><span class="vlist-r"><span class="vlist" style="height:0.2806em;"><span style="top:-2.55em;margin-right:0.05em;"><span class="pstrut" style="height:2.7em;"></span><span class="sizing reset-size6 size3 mtight"><span class="mord mathnormal mtight">t</span></span></span></span><span class="vlist-s">​</span></span><span class="vlist-r"><span class="vlist" style="height:0.15em;"><span></span></span></span></span></span></span><span class="mspace" style="margin-right:0.2778em;"></span><span class="mrel">=</span><span class="mspace" style="margin-right:0.2778em;"></span></span><span class="base"><span class="strut" style="height:1em;vertical-align:-0.25em;"></span><span class="mopen">(</span><span class="mord">1</span><span class="mspace" style="margin-right:0.2222em;"></span><span class="mbin">−</span><span class="mspace" style="margin-right:0.2222em;"></span></span><span class="base"><span class="strut" style="height:1em;vertical-align:-0.25em;"></span><span class="mord"><span class="mord mathnormal" style="margin-right:0.0528em;">β</span><span class="msupsub"><span class="vlist-t vlist-t2"><span class="vlist-r"><span class="vlist" style="height:0.2806em;"><span style="top:-2.55em;margin-left:-0.0528em;margin-right:0.05em;"><span class="pstrut" style="height:2.7em;"></span><span class="sizing reset-size6 size3 mtight"><span class="mord mathnormal mtight">t</span></span></span></span><span class="vlist-s">​</span></span><span class="vlist-r"><span class="vlist" style="height:0.15em;"><span></span></span></span></span></span></span><span class="mclose">)</span><span class="mspace" style="margin-right:0.2222em;"></span><span class="mbin">⋅</span><span class="mspace" style="margin-right:0.2222em;"></span></span><span class="base"><span class="strut" style="height:0.8917em;vertical-align:-0.2083em;"></span><span class="mord"><span class="mord text"><span class="mord">EMA</span></span><span class="msupsub"><span class="vlist-t vlist-t2"><span class="vlist-r"><span class="vlist" style="height:0.3011em;"><span style="top:-2.55em;margin-right:0.05em;"><span class="pstrut" style="height:2.7em;"></span><span class="sizing reset-size6 size3 mtight"><span class="mord mtight"><span class="mord mathnormal mtight">t</span><span class="mbin mtight">−</span><span class="mord mtight">1</span></span></span></span></span><span class="vlist-s">​</span></span><span class="vlist-r"><span class="vlist" style="height:0.2083em;"><span></span></span></span></span></span></span><span class="mspace" style="margin-right:0.2222em;"></span><span class="mbin">+</span><span class="mspace" style="margin-right:0.2222em;"></span></span><span class="base"><span class="strut" style="height:0.8889em;vertical-align:-0.1944em;"></span><span class="mord"><span class="mord mathnormal" style="margin-right:0.0528em;">β</span><span class="msupsub"><span class="vlist-t vlist-t2"><span class="vlist-r"><span class="vlist" style="height:0.2806em;"><span style="top:-2.55em;margin-left:-0.0528em;margin-right:0.05em;"><span class="pstrut" style="height:2.7em;"></span><span class="sizing reset-size6 size3 mtight"><span class="mord mathnormal mtight">t</span></span></span></span><span class="vlist-s">​</span></span><span class="vlist-r"><span class="vlist" style="height:0.15em;"><span></span></span></span></span></span></span><span class="mspace" style="margin-right:0.2222em;"></span><span class="mbin">⋅</span><span class="mspace" style="margin-right:0.2222em;"></span></span><span class="base"><span class="strut" style="height:0.8444em;vertical-align:-0.15em;"></span><span class="mord"><span class="mord mathnormal" style="margin-right:0.0278em;">θ</span><span class="msupsub"><span class="vlist-t vlist-t2"><span class="vlist-r"><span class="vlist" style="height:0.2806em;"><span style="top:-2.55em;margin-left:-0.0278em;margin-right:0.05em;"><span class="pstrut" style="height:2.7em;"></span><span class="sizing reset-size6 size3 mtight"><span class="mord mathnormal mtight">t</span></span></span></span><span class="vlist-s">​</span></span><span class="vlist-r"><span class="vlist" style="height:0.15em;"><span></span></span></span></span></span></span></span></span></span></span>`),n(d);var u=s(d,2),U=s(t(u));a(U,()=>'<span class="katex"><span class="katex-mathml"><math xmlns="http://www.w3.org/1998/Math/MathML"><semantics><mrow><msub><mi>β</mi><mi>t</mi></msub></mrow><annotation encoding="application/x-tex"> \\beta_t </annotation></semantics></math></span><span class="katex-html" aria-hidden="true"><span class="base"><span class="strut" style="height:0.8889em;vertical-align:-0.1944em;"></span><span class="mord"><span class="mord mathnormal" style="margin-right:0.0528em;">β</span><span class="msupsub"><span class="vlist-t vlist-t2"><span class="vlist-r"><span class="vlist" style="height:0.2806em;"><span style="top:-2.55em;margin-left:-0.0528em;margin-right:0.05em;"><span class="pstrut" style="height:2.7em;"></span><span class="sizing reset-size6 size3 mtight"><span class="mord mathnormal mtight">t</span></span></span></span><span class="vlist-s">​</span></span><span class="vlist-r"><span class="vlist" style="height:0.15em;"><span></span></span></span></span></span></span></span></span></span>');var F=s(U,2);a(F,()=>'<span class="katex"><span class="katex-mathml"><math xmlns="http://www.w3.org/1998/Math/MathML"><semantics><mrow><mi>t</mi></mrow><annotation encoding="application/x-tex"> t </annotation></semantics></math></span><span class="katex-html" aria-hidden="true"><span class="base"><span class="strut" style="height:0.6151em;"></span><span class="mord mathnormal">t</span></span></span></span>');var X=s(F,2);a(X,()=>`<span class="katex-display"><span class="katex"><span class="katex-mathml"><math xmlns="http://www.w3.org/1998/Math/MathML" display="block"><semantics><mrow><msub><mi>β</mi><mi>t</mi></msub><mo>=</mo><mo stretchy="false">(</mo><mi>ρ</mi><mo>+</mo><mi>γ</mi><mo>⋅</mo><mi>t</mi><msup><mo stretchy="false">)</mo><mrow><mo>−</mo><mi>κ</mi></mrow></msup><mi mathvariant="normal">.</mi></mrow><annotation encoding="application/x-tex">
\\beta_t = (\\rho + \\gamma \\cdot t)^{-\\kappa}.
</annotation></semantics></math></span><span class="katex-html" aria-hidden="true"><span class="base"><span class="strut" style="height:0.8889em;vertical-align:-0.1944em;"></span><span class="mord"><span class="mord mathnormal" style="margin-right:0.0528em;">β</span><span class="msupsub"><span class="vlist-t vlist-t2"><span class="vlist-r"><span class="vlist" style="height:0.2806em;"><span style="top:-2.55em;margin-left:-0.0528em;margin-right:0.05em;"><span class="pstrut" style="height:2.7em;"></span><span class="sizing reset-size6 size3 mtight"><span class="mord mathnormal mtight">t</span></span></span></span><span class="vlist-s">​</span></span><span class="vlist-r"><span class="vlist" style="height:0.15em;"><span></span></span></span></span></span></span><span class="mspace" style="margin-right:0.2778em;"></span><span class="mrel">=</span><span class="mspace" style="margin-right:0.2778em;"></span></span><span class="base"><span class="strut" style="height:1em;vertical-align:-0.25em;"></span><span class="mopen">(</span><span class="mord mathnormal">ρ</span><span class="mspace" style="margin-right:0.2222em;"></span><span class="mbin">+</span><span class="mspace" style="margin-right:0.2222em;"></span></span><span class="base"><span class="strut" style="height:0.6389em;vertical-align:-0.1944em;"></span><span class="mord mathnormal" style="margin-right:0.0556em;">γ</span><span class="mspace" style="margin-right:0.2222em;"></span><span class="mbin">⋅</span><span class="mspace" style="margin-right:0.2222em;"></span></span><span class="base"><span class="strut" style="height:1.0713em;vertical-align:-0.25em;"></span><span class="mord mathnormal">t</span><span class="mclose"><span class="mclose">)</span><span class="msupsub"><span class="vlist-t"><span class="vlist-r"><span class="vlist" style="height:0.8213em;"><span style="top:-3.113em;margin-right:0.05em;"><span class="pstrut" style="height:2.7em;"></span><span class="sizing reset-size6 size3 mtight"><span class="mord mtight"><span class="mord mtight">−</span><span class="mord mathnormal mtight">κ</span></span></span></span></span></span></span></span></span><span class="mord">.</span></span></span></span></span>`),n(u);var Y=s(u,2);ls(Y,{anchor:"trl.BEMACallback.example",children:(l,v)=>{var N=is(),H=s(L(N),2);j(H,{code:"ZnJvbSUyMHRybCUyMGltcG9ydCUyMEJFTUFDYWxsYmFjayUwQSUwQXRyYWluZXIlMjAlM0QlMjBUcmFpbmVyKC4uLiUyQyUyMGNhbGxiYWNrcyUzRCU1QkJFTUFDYWxsYmFjaygpJTVEKQ==",highlighted:`<span class="hljs-meta">&gt;&gt;&gt; </span><span class="hljs-keyword">from</span> trl <span class="hljs-keyword">import</span> BEMACallback
<span class="hljs-meta">&gt;&gt;&gt; </span>trainer = Trainer(..., callbacks=[BEMACallback()])`,lang:"python",wrap:!1}),y(l,N)},$$slots:{default:!0}}),n(c);var V=s(c,2);ss(V,{source:"https://github.com/huggingface/trl/blob/main/docs/source/bema_for_reference_model.md"}),m(2),y(W,f),es()}export{ds as component};

Xet Storage Details

Size:
39.1 kB
·
Xet hash:
d2d472fcee6c60d2eaf4ba0a649c9e86acb5fed1fae5656cfcbc80145846b6e0

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