Files
ReAgent/api/reagent.model_utils.html
2021-11-20 20:47:01 -08:00

280 lines
22 KiB
HTML
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
<!DOCTYPE html>
<html class="writer-html5" lang="en" >
<head>
<meta charset="utf-8" /><meta name="generator" content="Docutils 0.17.1: http://docutils.sourceforge.net/" />
<meta name="viewport" content="width=device-width, initial-scale=1.0" />
<title>reagent.model_utils package &mdash; ReAgent 1.0 documentation</title>
<link rel="stylesheet" href="../_static/pygments.css" type="text/css" />
<link rel="stylesheet" href="../_static/css/theme.css" type="text/css" />
<!--[if lt IE 9]>
<script src="../_static/js/html5shiv.min.js"></script>
<![endif]-->
<script data-url_root="../" id="documentation_options" src="../_static/documentation_options.js"></script>
<script src="../_static/jquery.js"></script>
<script src="../_static/underscore.js"></script>
<script src="../_static/doctools.js"></script>
<script src="../_static/js/theme.js"></script>
<link rel="index" title="Index" href="../genindex.html" />
<link rel="search" title="Search" href="../search.html" />
<link rel="next" title="reagent.net_builder package" href="reagent.net_builder.html" />
<link rel="prev" title="reagent.model_managers.ranking package" href="reagent.model_managers.ranking.html" />
</head>
<body class="wy-body-for-nav">
<div class="wy-grid-for-nav">
<nav data-toggle="wy-nav-shift" class="wy-nav-side">
<div class="wy-side-scroll">
<div class="wy-side-nav-search" >
<a href="../index.html" class="icon icon-home"> ReAgent
</a>
<div role="search">
<form id="rtd-search-form" class="wy-form" action="../search.html" method="get">
<input type="text" name="q" placeholder="Search docs" />
<input type="hidden" name="check_keywords" value="yes" />
<input type="hidden" name="area" value="default" />
</form>
</div>
</div><div class="wy-menu wy-menu-vertical" data-spy="affix" role="navigation" aria-label="Navigation menu">
<p class="caption" role="heading"><span class="caption-text">Getting Started</span></p>
<ul>
<li class="toctree-l1"><a class="reference internal" href="../installation.html">Installation</a></li>
<li class="toctree-l1"><a class="reference internal" href="../usage.html">Usage</a></li>
<li class="toctree-l1"><a class="reference internal" href="../rasp_tutorial.html">RASP (Not Actively Maintained)</a></li>
</ul>
<p class="caption" role="heading"><span class="caption-text">Advanced Topics</span></p>
<ul>
<li class="toctree-l1"><a class="reference internal" href="../continuous_integration.html">Continuous Integration</a></li>
</ul>
<p class="caption" role="heading"><span class="caption-text">Package Reference</span></p>
<ul class="current">
<li class="toctree-l1"><a class="reference internal" href="reagent.core.html">Core</a></li>
<li class="toctree-l1"><a class="reference internal" href="reagent.data.html">Data</a></li>
<li class="toctree-l1"><a class="reference internal" href="reagent.gym.html">Gym</a></li>
<li class="toctree-l1"><a class="reference internal" href="reagent.evaluation.html">Evaluation</a></li>
<li class="toctree-l1"><a class="reference internal" href="reagent.lite.html">Lite</a></li>
<li class="toctree-l1"><a class="reference internal" href="reagent.mab.html">MAB</a></li>
<li class="toctree-l1"><a class="reference internal" href="reagent.model_managers.html">Model Managers</a></li>
<li class="toctree-l1 current"><a class="current reference internal" href="#">Model Utils</a><ul>
<li class="toctree-l2"><a class="reference internal" href="#submodules">Submodules</a></li>
<li class="toctree-l2"><a class="reference internal" href="#module-reagent.model_utils.seq2slate_utils">reagent.model_utils.seq2slate_utils module</a></li>
<li class="toctree-l2"><a class="reference internal" href="#module-reagent.model_utils">Module contents</a></li>
</ul>
</li>
<li class="toctree-l1"><a class="reference internal" href="reagent.net_builder.html">Net Builders</a></li>
<li class="toctree-l1"><a class="reference internal" href="reagent.optimizer.html">Optimizers</a></li>
<li class="toctree-l1"><a class="reference internal" href="reagent.models.html">Models</a></li>
<li class="toctree-l1"><a class="reference internal" href="reagent.prediction.html">Prediction</a></li>
<li class="toctree-l1"><a class="reference internal" href="reagent.preprocessing.html">Preprocessing</a></li>
<li class="toctree-l1"><a class="reference internal" href="reagent.training.html">Training</a></li>
<li class="toctree-l1"><a class="reference internal" href="reagent.workflow.html">Workflow</a></li>
<li class="toctree-l1"><a class="reference internal" href="modules.html">All Modules</a></li>
</ul>
<p class="caption" role="heading"><span class="caption-text">Others</span></p>
<ul>
<li class="toctree-l1"><a class="reference external" href="https://github.com/facebookresearch/ReAgent">Github</a></li>
<li class="toctree-l1"><a class="reference internal" href="../license.html">License</a></li>
</ul>
</div>
</div>
</nav>
<section data-toggle="wy-nav-shift" class="wy-nav-content-wrap"><nav class="wy-nav-top" aria-label="Mobile navigation menu" >
<i data-toggle="wy-nav-top" class="fa fa-bars"></i>
<a href="../index.html">ReAgent</a>
</nav>
<div class="wy-nav-content">
<div class="rst-content">
<div role="navigation" aria-label="Page navigation">
<ul class="wy-breadcrumbs">
<li><a href="../index.html" class="icon icon-home"></a> &raquo;</li>
<li>reagent.model_utils package</li>
<li class="wy-breadcrumbs-aside">
<a href="../_sources/api/reagent.model_utils.rst.txt" rel="nofollow"> View page source</a>
</li>
</ul>
<hr/>
</div>
<div role="main" class="document" itemscope="itemscope" itemtype="http://schema.org/Article">
<div itemprop="articleBody">
<section id="reagent-model-utils-package">
<h1>reagent.model_utils package<a class="headerlink" href="#reagent-model-utils-package" title="Permalink to this headline"></a></h1>
<section id="submodules">
<h2>Submodules<a class="headerlink" href="#submodules" title="Permalink to this headline"></a></h2>
</section>
<section id="module-reagent.model_utils.seq2slate_utils">
<span id="reagent-model-utils-seq2slate-utils-module"></span><h2>reagent.model_utils.seq2slate_utils module<a class="headerlink" href="#module-reagent.model_utils.seq2slate_utils" title="Permalink to this headline"></a></h2>
<dl class="py class">
<dt class="sig sig-object py" id="reagent.model_utils.seq2slate_utils.Seq2SlateMode">
<em class="property"><span class="pre">class</span><span class="w"> </span></em><span class="sig-prename descclassname"><span class="pre">reagent.model_utils.seq2slate_utils.</span></span><span class="sig-name descname"><span class="pre">Seq2SlateMode</span></span><span class="sig-paren">(</span><em class="sig-param"><span class="n"><span class="pre">value</span></span></em><span class="sig-paren">)</span><a class="headerlink" href="#reagent.model_utils.seq2slate_utils.Seq2SlateMode" title="Permalink to this definition"></a></dt>
<dd><p>Bases: <code class="xref py py-class docutils literal notranslate"><span class="pre">enum.Enum</span></code></p>
<p>An enumeration.</p>
<dl class="py attribute">
<dt class="sig sig-object py" id="reagent.model_utils.seq2slate_utils.Seq2SlateMode.DECODE_ONE_STEP_MODE">
<span class="sig-name descname"><span class="pre">DECODE_ONE_STEP_MODE</span></span><em class="property"><span class="w"> </span><span class="p"><span class="pre">=</span></span><span class="w"> </span><span class="pre">'decode_one_step'</span></em><a class="headerlink" href="#reagent.model_utils.seq2slate_utils.Seq2SlateMode.DECODE_ONE_STEP_MODE" title="Permalink to this definition"></a></dt>
<dd></dd></dl>
<dl class="py attribute">
<dt class="sig sig-object py" id="reagent.model_utils.seq2slate_utils.Seq2SlateMode.ENCODER_SCORE_MODE">
<span class="sig-name descname"><span class="pre">ENCODER_SCORE_MODE</span></span><em class="property"><span class="w"> </span><span class="p"><span class="pre">=</span></span><span class="w"> </span><span class="pre">'encoder_score_mode'</span></em><a class="headerlink" href="#reagent.model_utils.seq2slate_utils.Seq2SlateMode.ENCODER_SCORE_MODE" title="Permalink to this definition"></a></dt>
<dd></dd></dl>
<dl class="py attribute">
<dt class="sig sig-object py" id="reagent.model_utils.seq2slate_utils.Seq2SlateMode.PER_SEQ_LOG_PROB_MODE">
<span class="sig-name descname"><span class="pre">PER_SEQ_LOG_PROB_MODE</span></span><em class="property"><span class="w"> </span><span class="p"><span class="pre">=</span></span><span class="w"> </span><span class="pre">'per_sequence_log_prob'</span></em><a class="headerlink" href="#reagent.model_utils.seq2slate_utils.Seq2SlateMode.PER_SEQ_LOG_PROB_MODE" title="Permalink to this definition"></a></dt>
<dd></dd></dl>
<dl class="py attribute">
<dt class="sig sig-object py" id="reagent.model_utils.seq2slate_utils.Seq2SlateMode.PER_SYMBOL_LOG_PROB_DIST_MODE">
<span class="sig-name descname"><span class="pre">PER_SYMBOL_LOG_PROB_DIST_MODE</span></span><em class="property"><span class="w"> </span><span class="p"><span class="pre">=</span></span><span class="w"> </span><span class="pre">'per_symbol_log_prob_dist'</span></em><a class="headerlink" href="#reagent.model_utils.seq2slate_utils.Seq2SlateMode.PER_SYMBOL_LOG_PROB_DIST_MODE" title="Permalink to this definition"></a></dt>
<dd></dd></dl>
<dl class="py attribute">
<dt class="sig sig-object py" id="reagent.model_utils.seq2slate_utils.Seq2SlateMode.RANK_MODE">
<span class="sig-name descname"><span class="pre">RANK_MODE</span></span><em class="property"><span class="w"> </span><span class="p"><span class="pre">=</span></span><span class="w"> </span><span class="pre">'rank'</span></em><a class="headerlink" href="#reagent.model_utils.seq2slate_utils.Seq2SlateMode.RANK_MODE" title="Permalink to this definition"></a></dt>
<dd></dd></dl>
</dd></dl>
<dl class="py class">
<dt class="sig sig-object py" id="reagent.model_utils.seq2slate_utils.Seq2SlateOutputArch">
<em class="property"><span class="pre">class</span><span class="w"> </span></em><span class="sig-prename descclassname"><span class="pre">reagent.model_utils.seq2slate_utils.</span></span><span class="sig-name descname"><span class="pre">Seq2SlateOutputArch</span></span><span class="sig-paren">(</span><em class="sig-param"><span class="n"><span class="pre">value</span></span></em><span class="sig-paren">)</span><a class="headerlink" href="#reagent.model_utils.seq2slate_utils.Seq2SlateOutputArch" title="Permalink to this definition"></a></dt>
<dd><p>Bases: <code class="xref py py-class docutils literal notranslate"><span class="pre">enum.Enum</span></code></p>
<p>An enumeration.</p>
<dl class="py attribute">
<dt class="sig sig-object py" id="reagent.model_utils.seq2slate_utils.Seq2SlateOutputArch.AUTOREGRESSIVE">
<span class="sig-name descname"><span class="pre">AUTOREGRESSIVE</span></span><em class="property"><span class="w"> </span><span class="p"><span class="pre">=</span></span><span class="w"> </span><span class="pre">'autoregressive'</span></em><a class="headerlink" href="#reagent.model_utils.seq2slate_utils.Seq2SlateOutputArch.AUTOREGRESSIVE" title="Permalink to this definition"></a></dt>
<dd></dd></dl>
<dl class="py attribute">
<dt class="sig sig-object py" id="reagent.model_utils.seq2slate_utils.Seq2SlateOutputArch.ENCODER_SCORE">
<span class="sig-name descname"><span class="pre">ENCODER_SCORE</span></span><em class="property"><span class="w"> </span><span class="p"><span class="pre">=</span></span><span class="w"> </span><span class="pre">'encoder_score'</span></em><a class="headerlink" href="#reagent.model_utils.seq2slate_utils.Seq2SlateOutputArch.ENCODER_SCORE" title="Permalink to this definition"></a></dt>
<dd></dd></dl>
<dl class="py attribute">
<dt class="sig sig-object py" id="reagent.model_utils.seq2slate_utils.Seq2SlateOutputArch.FRECHET_SORT">
<span class="sig-name descname"><span class="pre">FRECHET_SORT</span></span><em class="property"><span class="w"> </span><span class="p"><span class="pre">=</span></span><span class="w"> </span><span class="pre">'frechet_sort'</span></em><a class="headerlink" href="#reagent.model_utils.seq2slate_utils.Seq2SlateOutputArch.FRECHET_SORT" title="Permalink to this definition"></a></dt>
<dd></dd></dl>
</dd></dl>
<dl class="py function">
<dt class="sig sig-object py" id="reagent.model_utils.seq2slate_utils.attention">
<span class="sig-prename descclassname"><span class="pre">reagent.model_utils.seq2slate_utils.</span></span><span class="sig-name descname"><span class="pre">attention</span></span><span class="sig-paren">(</span><em class="sig-param"><span class="n"><span class="pre">query</span></span></em>, <em class="sig-param"><span class="n"><span class="pre">key</span></span></em>, <em class="sig-param"><span class="n"><span class="pre">value</span></span></em>, <em class="sig-param"><span class="n"><span class="pre">mask</span></span></em>, <em class="sig-param"><span class="n"><span class="pre">d_k</span></span></em><span class="sig-paren">)</span><a class="headerlink" href="#reagent.model_utils.seq2slate_utils.attention" title="Permalink to this definition"></a></dt>
<dd><p>Scaled Dot Product Attention</p>
</dd></dl>
<dl class="py function">
<dt class="sig sig-object py" id="reagent.model_utils.seq2slate_utils.clones">
<span class="sig-prename descclassname"><span class="pre">reagent.model_utils.seq2slate_utils.</span></span><span class="sig-name descname"><span class="pre">clones</span></span><span class="sig-paren">(</span><em class="sig-param"><span class="n"><span class="pre">module</span></span></em>, <em class="sig-param"><span class="n"><span class="pre">N</span></span></em><span class="sig-paren">)</span><a class="headerlink" href="#reagent.model_utils.seq2slate_utils.clones" title="Permalink to this definition"></a></dt>
<dd><p>Produce N identical layers.</p>
<dl class="field-list simple">
<dt class="field-odd">Parameters</dt>
<dd class="field-odd"><ul class="simple">
<li><p><strong>module</strong> – nn.Module class</p></li>
<li><p><strong>N</strong> – number of copies</p></li>
</ul>
</dd>
</dl>
</dd></dl>
<dl class="py function">
<dt class="sig sig-object py" id="reagent.model_utils.seq2slate_utils.mask_logits_by_idx">
<span class="sig-prename descclassname"><span class="pre">reagent.model_utils.seq2slate_utils.</span></span><span class="sig-name descname"><span class="pre">mask_logits_by_idx</span></span><span class="sig-paren">(</span><em class="sig-param"><span class="n"><span class="pre">logits</span></span></em>, <em class="sig-param"><span class="n"><span class="pre">tgt_in_idx</span></span></em><span class="sig-paren">)</span><a class="headerlink" href="#reagent.model_utils.seq2slate_utils.mask_logits_by_idx" title="Permalink to this definition"></a></dt>
<dd></dd></dl>
<dl class="py function">
<dt class="sig sig-object py" id="reagent.model_utils.seq2slate_utils.per_symbol_to_per_seq_log_probs">
<span class="sig-prename descclassname"><span class="pre">reagent.model_utils.seq2slate_utils.</span></span><span class="sig-name descname"><span class="pre">per_symbol_to_per_seq_log_probs</span></span><span class="sig-paren">(</span><em class="sig-param"><span class="n"><span class="pre">per_symbol_log_probs</span></span></em>, <em class="sig-param"><span class="n"><span class="pre">tgt_out_idx</span></span></em><span class="sig-paren">)</span><a class="headerlink" href="#reagent.model_utils.seq2slate_utils.per_symbol_to_per_seq_log_probs" title="Permalink to this definition"></a></dt>
<dd><p>Gather per-symbol log probabilities into per-seq log probabilities</p>
</dd></dl>
<dl class="py function">
<dt class="sig sig-object py" id="reagent.model_utils.seq2slate_utils.per_symbol_to_per_seq_probs">
<span class="sig-prename descclassname"><span class="pre">reagent.model_utils.seq2slate_utils.</span></span><span class="sig-name descname"><span class="pre">per_symbol_to_per_seq_probs</span></span><span class="sig-paren">(</span><em class="sig-param"><span class="n"><span class="pre">per_symbol_probs</span></span></em>, <em class="sig-param"><span class="n"><span class="pre">tgt_out_idx</span></span></em><span class="sig-paren">)</span><a class="headerlink" href="#reagent.model_utils.seq2slate_utils.per_symbol_to_per_seq_probs" title="Permalink to this definition"></a></dt>
<dd><p>Gather per-symbol probabilities into per-seq probabilities</p>
</dd></dl>
<dl class="py function">
<dt class="sig sig-object py" id="reagent.model_utils.seq2slate_utils.print_model_info">
<span class="sig-prename descclassname"><span class="pre">reagent.model_utils.seq2slate_utils.</span></span><span class="sig-name descname"><span class="pre">print_model_info</span></span><span class="sig-paren">(</span><em class="sig-param"><span class="n"><span class="pre">seq2slate</span></span></em><span class="sig-paren">)</span><a class="headerlink" href="#reagent.model_utils.seq2slate_utils.print_model_info" title="Permalink to this definition"></a></dt>
<dd></dd></dl>
<dl class="py function">
<dt class="sig sig-object py" id="reagent.model_utils.seq2slate_utils.pytorch_decoder_mask">
<span class="sig-prename descclassname"><span class="pre">reagent.model_utils.seq2slate_utils.</span></span><span class="sig-name descname"><span class="pre">pytorch_decoder_mask</span></span><span class="sig-paren">(</span><em class="sig-param"><span class="n"><span class="pre">memory</span></span><span class="p"><span class="pre">:</span></span><span class="w"> </span><span class="n"><span class="pre">torch.Tensor</span></span></em>, <em class="sig-param"><span class="n"><span class="pre">tgt_in_idx</span></span><span class="p"><span class="pre">:</span></span><span class="w"> </span><span class="n"><span class="pre">torch.Tensor</span></span></em>, <em class="sig-param"><span class="n"><span class="pre">num_heads</span></span><span class="p"><span class="pre">:</span></span><span class="w"> </span><span class="n"><span class="pre">int</span></span></em><span class="sig-paren">)</span><a class="headerlink" href="#reagent.model_utils.seq2slate_utils.pytorch_decoder_mask" title="Permalink to this definition"></a></dt>
<dd><p>Compute the masks used in the PyTorch Transformer-based decoder for
self-attention and attention over encoder outputs</p>
<p>mask_ijk = 1 if the item should be ignored; 0 if the item should be paid attention</p>
<dl class="simple">
<dt>Input:</dt><dd><p>memory shape: batch_size, src_seq_len, dim_model
tgt_in_idx (+2 offseted) shape: batch_size, tgt_seq_len</p>
</dd>
</dl>
<dl class="field-list simple">
<dt class="field-odd">Returns</dt>
<dd class="field-odd"><p>batch_size * num_heads, tgt_seq_len, tgt_seq_len
tgt_src_mask shape: batch_size * num_heads, tgt_seq_len, src_seq_len</p>
</dd>
<dt class="field-even">Return type</dt>
<dd class="field-even"><p>tgt_tgt_mask shape</p>
</dd>
</dl>
</dd></dl>
<dl class="py function">
<dt class="sig sig-object py" id="reagent.model_utils.seq2slate_utils.subsequent_and_padding_mask">
<span class="sig-prename descclassname"><span class="pre">reagent.model_utils.seq2slate_utils.</span></span><span class="sig-name descname"><span class="pre">subsequent_and_padding_mask</span></span><span class="sig-paren">(</span><em class="sig-param"><span class="n"><span class="pre">tgt_in_idx</span></span></em><span class="sig-paren">)</span><a class="headerlink" href="#reagent.model_utils.seq2slate_utils.subsequent_and_padding_mask" title="Permalink to this definition"></a></dt>
<dd><p>Create a mask to hide padding and future items</p>
</dd></dl>
<dl class="py function">
<dt class="sig sig-object py" id="reagent.model_utils.seq2slate_utils.subsequent_mask">
<span class="sig-prename descclassname"><span class="pre">reagent.model_utils.seq2slate_utils.</span></span><span class="sig-name descname"><span class="pre">subsequent_mask</span></span><span class="sig-paren">(</span><em class="sig-param"><span class="n"><span class="pre">size</span></span><span class="p"><span class="pre">:</span></span><span class="w"> </span><span class="n"><span class="pre">int</span></span></em>, <em class="sig-param"><span class="n"><span class="pre">device</span></span><span class="p"><span class="pre">:</span></span><span class="w"> </span><span class="n"><span class="pre">torch.device</span></span></em><span class="sig-paren">)</span><a class="headerlink" href="#reagent.model_utils.seq2slate_utils.subsequent_mask" title="Permalink to this definition"></a></dt>
<dd><p>Mask out subsequent positions. Mainly used in the decoding process,
in which an item should not attend subsequent items.</p>
<p>mask_ijk = 0 if the item should be ignored; 1 if the item should be paid attention</p>
</dd></dl>
</section>
<section id="module-reagent.model_utils">
<span id="module-contents"></span><h2>Module contents<a class="headerlink" href="#module-reagent.model_utils" title="Permalink to this headline"></a></h2>
</section>
</section>
</div>
</div>
<footer><div class="rst-footer-buttons" role="navigation" aria-label="Footer">
<a href="reagent.model_managers.ranking.html" class="btn btn-neutral float-left" title="reagent.model_managers.ranking package" accesskey="p" rel="prev"><span class="fa fa-arrow-circle-left" aria-hidden="true"></span> Previous</a>
<a href="reagent.net_builder.html" class="btn btn-neutral float-right" title="reagent.net_builder package" accesskey="n" rel="next">Next <span class="fa fa-arrow-circle-right" aria-hidden="true"></span></a>
</div>
<hr/>
<div role="contentinfo">
<p>&#169; Copyright 2022, Meta Platforms, Inc.</p>
</div>
Built with <a href="https://www.sphinx-doc.org/">Sphinx</a> using a
<a href="https://github.com/readthedocs/sphinx_rtd_theme">theme</a>
provided by <a href="https://readthedocs.org">Read the Docs</a>.
</footer>
</div>
</div>
</section>
</div>
<script>
jQuery(function () {
SphinxRtdTheme.Navigation.enable(true);
});
</script>
</body>
</html>