facebookresearch--audiocraft
858 行
52 KiB
HTML
858 行
52 KiB
HTML
<!doctype html>
|
||
<html lang="en">
|
||
<head>
|
||
<meta charset="utf-8">
|
||
<meta name="viewport" content="width=device-width, initial-scale=1, minimum-scale=1">
|
||
<meta name="generator" content="pdoc3 0.11.5">
|
||
<title>audiocraft.utils.samples.manager API documentation</title>
|
||
<meta name="description" content="API that can manage the storage and retrieval of generated samples produced by experiments …">
|
||
<link rel="stylesheet" href="https://cdnjs.cloudflare.com/ajax/libs/10up-sanitize.css/13.0.0/sanitize.min.css" integrity="sha512-y1dtMcuvtTMJc1yPgEqF0ZjQbhnc/bFhyvIyVNb9Zk5mIGtqVaAB1Ttl28su8AvFMOY0EwRbAe+HCLqj6W7/KA==" crossorigin>
|
||
<link rel="stylesheet" href="https://cdnjs.cloudflare.com/ajax/libs/10up-sanitize.css/13.0.0/typography.min.css" integrity="sha512-Y1DYSb995BAfxobCkKepB1BqJJTPrOp3zPL74AWFugHHmmdcvO+C48WLrUOlhGMc0QG7AE3f7gmvvcrmX2fDoA==" crossorigin>
|
||
<link rel="stylesheet" href="https://cdnjs.cloudflare.com/ajax/libs/highlight.js/11.9.0/styles/default.min.css" crossorigin>
|
||
<style>:root{--highlight-color:#fe9}.flex{display:flex !important}body{line-height:1.5em}#content{padding:20px}#sidebar{padding:1.5em;overflow:hidden}#sidebar > *:last-child{margin-bottom:2cm}.http-server-breadcrumbs{font-size:130%;margin:0 0 15px 0}#footer{font-size:.75em;padding:5px 30px;border-top:1px solid #ddd;text-align:right}#footer p{margin:0 0 0 1em;display:inline-block}#footer p:last-child{margin-right:30px}h1,h2,h3,h4,h5{font-weight:300}h1{font-size:2.5em;line-height:1.1em}h2{font-size:1.75em;margin:2em 0 .50em 0}h3{font-size:1.4em;margin:1.6em 0 .7em 0}h4{margin:0;font-size:105%}h1:target,h2:target,h3:target,h4:target,h5:target,h6:target{background:var(--highlight-color);padding:.2em 0}a{color:#058;text-decoration:none;transition:color .2s ease-in-out}a:visited{color:#503}a:hover{color:#b62}.title code{font-weight:bold}h2[id^="header-"]{margin-top:2em}.ident{color:#900;font-weight:bold}pre code{font-size:.8em;line-height:1.4em;padding:1em;display:block}code{background:#f3f3f3;font-family:"DejaVu Sans Mono",monospace;padding:1px 4px;overflow-wrap:break-word}h1 code{background:transparent}pre{border-top:1px solid #ccc;border-bottom:1px solid #ccc;margin:1em 0}#http-server-module-list{display:flex;flex-flow:column}#http-server-module-list div{display:flex}#http-server-module-list dt{min-width:10%}#http-server-module-list p{margin-top:0}.toc ul,#index{list-style-type:none;margin:0;padding:0}#index code{background:transparent}#index h3{border-bottom:1px solid #ddd}#index ul{padding:0}#index h4{margin-top:.6em;font-weight:bold}@media (min-width:200ex){#index .two-column{column-count:2}}@media (min-width:300ex){#index .two-column{column-count:3}}dl{margin-bottom:2em}dl dl:last-child{margin-bottom:4em}dd{margin:0 0 1em 3em}#header-classes + dl > dd{margin-bottom:3em}dd dd{margin-left:2em}dd p{margin:10px 0}.name{background:#eee;font-size:.85em;padding:5px 10px;display:inline-block;min-width:40%}.name:hover{background:#e0e0e0}dt:target .name{background:var(--highlight-color)}.name > span:first-child{white-space:nowrap}.name.class > span:nth-child(2){margin-left:.4em}.inherited{color:#999;border-left:5px solid #eee;padding-left:1em}.inheritance em{font-style:normal;font-weight:bold}.desc h2{font-weight:400;font-size:1.25em}.desc h3{font-size:1em}.desc dt code{background:inherit}.source > summary,.git-link-div{color:#666;text-align:right;font-weight:400;font-size:.8em;text-transform:uppercase}.source summary > *{white-space:nowrap;cursor:pointer}.git-link{color:inherit;margin-left:1em}.source pre{max-height:500px;overflow:auto;margin:0}.source pre code{font-size:12px;overflow:visible;min-width:max-content}.hlist{list-style:none}.hlist li{display:inline}.hlist li:after{content:',\2002'}.hlist li:last-child:after{content:none}.hlist .hlist{display:inline;padding-left:1em}img{max-width:100%}td{padding:0 .5em}.admonition{padding:.1em 1em;margin:1em 0}.admonition-title{font-weight:bold}.admonition.note,.admonition.info,.admonition.important{background:#aef}.admonition.todo,.admonition.versionadded,.admonition.tip,.admonition.hint{background:#dfd}.admonition.warning,.admonition.versionchanged,.admonition.deprecated{background:#fd4}.admonition.error,.admonition.danger,.admonition.caution{background:lightpink}</style>
|
||
<style media="screen and (min-width: 700px)">@media screen and (min-width:700px){#sidebar{width:30%;height:100vh;overflow:auto;position:sticky;top:0}#content{width:70%;max-width:100ch;padding:3em 4em;border-left:1px solid #ddd}pre code{font-size:1em}.name{font-size:1em}main{display:flex;flex-direction:row-reverse;justify-content:flex-end}.toc ul ul,#index ul ul{padding-left:1em}.toc > ul > li{margin-top:.5em}}</style>
|
||
<style media="print">@media print{#sidebar h1{page-break-before:always}.source{display:none}}@media print{*{background:transparent !important;color:#000 !important;box-shadow:none !important;text-shadow:none !important}a[href]:after{content:" (" attr(href) ")";font-size:90%}a[href][title]:after{content:none}abbr[title]:after{content:" (" attr(title) ")"}.ir a:after,a[href^="javascript:"]:after,a[href^="#"]:after{content:""}pre,blockquote{border:1px solid #999;page-break-inside:avoid}thead{display:table-header-group}tr,img{page-break-inside:avoid}img{max-width:100% !important}@page{margin:0.5cm}p,h2,h3{orphans:3;widows:3}h1,h2,h3,h4,h5,h6{page-break-after:avoid}}</style>
|
||
<script defer src="https://cdnjs.cloudflare.com/ajax/libs/highlight.js/11.9.0/highlight.min.js" integrity="sha512-D9gUyxqja7hBtkWpPWGt9wfbfaMGVt9gnyCvYa+jojwwPHLCzUm5i8rpk7vD7wNee9bA35eYIjobYPaQuKS1MQ==" crossorigin></script>
|
||
<script>window.addEventListener('DOMContentLoaded', () => {
|
||
hljs.configure({languages: ['bash', 'css', 'diff', 'graphql', 'ini', 'javascript', 'json', 'plaintext', 'python', 'python-repl', 'rust', 'shell', 'sql', 'typescript', 'xml', 'yaml']});
|
||
hljs.highlightAll();
|
||
/* Collapse source docstrings */
|
||
setTimeout(() => {
|
||
[...document.querySelectorAll('.hljs.language-python > .hljs-string')]
|
||
.filter(el => el.innerHTML.length > 200 && ['"""', "'''"].includes(el.innerHTML.substring(0, 3)))
|
||
.forEach(el => {
|
||
let d = document.createElement('details');
|
||
d.classList.add('hljs-string');
|
||
d.innerHTML = '<summary>"""</summary>' + el.innerHTML.substring(3);
|
||
el.replaceWith(d);
|
||
});
|
||
}, 100);
|
||
})</script>
|
||
</head>
|
||
<body>
|
||
<main>
|
||
<article id="content">
|
||
<header>
|
||
<h1 class="title">Module <code>audiocraft.utils.samples.manager</code></h1>
|
||
</header>
|
||
<section id="section-intro">
|
||
<p>API that can manage the storage and retrieval of generated samples produced by experiments.</p>
|
||
<p>It offers the following benefits:
|
||
* Samples are stored in a consistent way across epoch
|
||
* Metadata about the samples can be stored and retrieved
|
||
* Can retrieve audio
|
||
* Identifiers are reliable and deterministic for prompted and conditioned samples
|
||
* Can request the samples for multiple XPs, grouped by sample identifier
|
||
* For no-input samples (not prompt and no conditions), samples across XPs are matched
|
||
by sorting their identifiers</p>
|
||
</section>
|
||
<section>
|
||
</section>
|
||
<section>
|
||
</section>
|
||
<section>
|
||
<h2 class="section-title" id="header-functions">Functions</h2>
|
||
<dl>
|
||
<dt id="audiocraft.utils.samples.manager.get_samples_for_xps"><code class="name flex">
|
||
<span>def <span class="ident">get_samples_for_xps</span></span>(<span>xps: List[dora.xp.XP], **kwargs) ‑> Dict[str, List[<a title="audiocraft.utils.samples.manager.Sample" href="#audiocraft.utils.samples.manager.Sample">Sample</a>]]</span>
|
||
</code></dt>
|
||
<dd>
|
||
<details class="source">
|
||
<summary>
|
||
<span>Expand source code</span>
|
||
</summary>
|
||
<pre><code class="python">def get_samples_for_xps(xps: tp.List[dora.XP], **kwargs) -> tp.Dict[str, tp.List[Sample]]:
|
||
"""Gets a dictionary of matched samples across the given XPs.
|
||
Each dictionary entry maps a sample id to a list of samples for that id. The number of samples per id
|
||
will always match the number of XPs provided and will correspond to each XP in the same order given.
|
||
In other words, only samples that can be match across all provided XPs will be returned
|
||
in order to satisfy this rule.
|
||
|
||
There are two types of ids that can be returned: stable and unstable.
|
||
* Stable IDs are deterministic ids that were computed by the SampleManager given a sample's inputs
|
||
(prompts/conditioning). This is why we can match them across XPs.
|
||
* Unstable IDs are of the form "noinput_{idx}" and are generated on-the-fly, in order to map samples
|
||
that used non-deterministic, random ids. This is the case for samples that did not use prompts or
|
||
conditioning for their generation. This function will sort these samples by their id and match them
|
||
by their index.
|
||
|
||
Args:
|
||
xps: a list of XPs to match samples from.
|
||
start_epoch (int): If provided, only return samples corresponding to this epoch or newer.
|
||
end_epoch (int): If provided, only return samples corresponding to this epoch or older.
|
||
exclude_prompted (bool): If True, does not include samples that used a prompt.
|
||
exclude_unprompted (bool): If True, does not include samples that did not use a prompt.
|
||
exclude_conditioned (bool): If True, excludes samples that used conditioning.
|
||
exclude_unconditioned (bool): If True, excludes samples that did not use conditioning.
|
||
"""
|
||
managers = [SampleManager(xp) for xp in xps]
|
||
samples_per_xp = [manager.get_samples(**kwargs) for manager in managers]
|
||
stable_samples = _match_stable_samples(samples_per_xp)
|
||
unstable_samples = _match_unstable_samples(samples_per_xp)
|
||
return dict(stable_samples, **unstable_samples)</code></pre>
|
||
</details>
|
||
<div class="desc"><p>Gets a dictionary of matched samples across the given XPs.
|
||
Each dictionary entry maps a sample id to a list of samples for that id. The number of samples per id
|
||
will always match the number of XPs provided and will correspond to each XP in the same order given.
|
||
In other words, only samples that can be match across all provided XPs will be returned
|
||
in order to satisfy this rule.</p>
|
||
<p>There are two types of ids that can be returned: stable and unstable.
|
||
* Stable IDs are deterministic ids that were computed by the SampleManager given a sample's inputs
|
||
(prompts/conditioning). This is why we can match them across XPs.
|
||
* Unstable IDs are of the form "noinput_{idx}" and are generated on-the-fly, in order to map samples
|
||
that used non-deterministic, random ids. This is the case for samples that did not use prompts or
|
||
conditioning for their generation. This function will sort these samples by their id and match them
|
||
by their index.</p>
|
||
<h2 id="args">Args</h2>
|
||
<dl>
|
||
<dt><strong><code>xps</code></strong></dt>
|
||
<dd>a list of XPs to match samples from.</dd>
|
||
<dt><strong><code>start_epoch</code></strong> : <code>int</code></dt>
|
||
<dd>If provided, only return samples corresponding to this epoch or newer.</dd>
|
||
<dt><strong><code>end_epoch</code></strong> : <code>int</code></dt>
|
||
<dd>If provided, only return samples corresponding to this epoch or older.</dd>
|
||
<dt><strong><code>exclude_prompted</code></strong> : <code>bool</code></dt>
|
||
<dd>If True, does not include samples that used a prompt.</dd>
|
||
<dt><strong><code>exclude_unprompted</code></strong> : <code>bool</code></dt>
|
||
<dd>If True, does not include samples that did not use a prompt.</dd>
|
||
<dt><strong><code>exclude_conditioned</code></strong> : <code>bool</code></dt>
|
||
<dd>If True, excludes samples that used conditioning.</dd>
|
||
<dt><strong><code>exclude_unconditioned</code></strong> : <code>bool</code></dt>
|
||
<dd>If True, excludes samples that did not use conditioning.</dd>
|
||
</dl></div>
|
||
</dd>
|
||
<dt id="audiocraft.utils.samples.manager.slugify"><code class="name flex">
|
||
<span>def <span class="ident">slugify</span></span>(<span>value: Any, allow_unicode: bool = False)</span>
|
||
</code></dt>
|
||
<dd>
|
||
<details class="source">
|
||
<summary>
|
||
<span>Expand source code</span>
|
||
</summary>
|
||
<pre><code class="python">def slugify(value: tp.Any, allow_unicode: bool = False):
|
||
"""Process string for safer file naming.
|
||
|
||
Taken from https://github.com/django/django/blob/master/django/utils/text.py
|
||
|
||
Convert to ASCII if 'allow_unicode' is False. Convert spaces or repeated
|
||
dashes to single dashes. Remove characters that aren't alphanumerics,
|
||
underscores, or hyphens. Convert to lowercase. Also strip leading and
|
||
trailing whitespace, dashes, and underscores.
|
||
"""
|
||
value = str(value)
|
||
if allow_unicode:
|
||
value = unicodedata.normalize("NFKC", value)
|
||
else:
|
||
value = (
|
||
unicodedata.normalize("NFKD", value)
|
||
.encode("ascii", "ignore")
|
||
.decode("ascii")
|
||
)
|
||
value = re.sub(r"[^\w\s-]", "", value.lower())
|
||
return re.sub(r"[-\s]+", "-", value).strip("-_")</code></pre>
|
||
</details>
|
||
<div class="desc"><p>Process string for safer file naming.</p>
|
||
<p>Taken from <a href="https://github.com/django/django/blob/master/django/utils/text.py">https://github.com/django/django/blob/master/django/utils/text.py</a></p>
|
||
<p>Convert to ASCII if 'allow_unicode' is False. Convert spaces or repeated
|
||
dashes to single dashes. Remove characters that aren't alphanumerics,
|
||
underscores, or hyphens. Convert to lowercase. Also strip leading and
|
||
trailing whitespace, dashes, and underscores.</p></div>
|
||
</dd>
|
||
</dl>
|
||
</section>
|
||
<section>
|
||
<h2 class="section-title" id="header-classes">Classes</h2>
|
||
<dl>
|
||
<dt id="audiocraft.utils.samples.manager.ReferenceSample"><code class="flex name class">
|
||
<span>class <span class="ident">ReferenceSample</span></span>
|
||
<span>(</span><span>id: str, path: str, duration: float)</span>
|
||
</code></dt>
|
||
<dd>
|
||
<details class="source">
|
||
<summary>
|
||
<span>Expand source code</span>
|
||
</summary>
|
||
<pre><code class="python">@dataclass
|
||
class ReferenceSample:
|
||
id: str
|
||
path: str
|
||
duration: float</code></pre>
|
||
</details>
|
||
<div class="desc"><p>ReferenceSample(id: str, path: str, duration: float)</p></div>
|
||
<h3>Class variables</h3>
|
||
<dl>
|
||
<dt id="audiocraft.utils.samples.manager.ReferenceSample.duration"><code class="name">var <span class="ident">duration</span> : float</code></dt>
|
||
<dd>
|
||
<div class="desc"></div>
|
||
</dd>
|
||
<dt id="audiocraft.utils.samples.manager.ReferenceSample.id"><code class="name">var <span class="ident">id</span> : str</code></dt>
|
||
<dd>
|
||
<div class="desc"></div>
|
||
</dd>
|
||
<dt id="audiocraft.utils.samples.manager.ReferenceSample.path"><code class="name">var <span class="ident">path</span> : str</code></dt>
|
||
<dd>
|
||
<div class="desc"></div>
|
||
</dd>
|
||
</dl>
|
||
</dd>
|
||
<dt id="audiocraft.utils.samples.manager.Sample"><code class="flex name class">
|
||
<span>class <span class="ident">Sample</span></span>
|
||
<span>(</span><span>id: str,<br>path: str,<br>epoch: int,<br>duration: float,<br>conditioning: Dict[str, Any] | None,<br>prompt: <a title="audiocraft.utils.samples.manager.ReferenceSample" href="#audiocraft.utils.samples.manager.ReferenceSample">ReferenceSample</a> | None,<br>reference: <a title="audiocraft.utils.samples.manager.ReferenceSample" href="#audiocraft.utils.samples.manager.ReferenceSample">ReferenceSample</a> | None,<br>generation_args: Dict[str, Any] | None)</span>
|
||
</code></dt>
|
||
<dd>
|
||
<details class="source">
|
||
<summary>
|
||
<span>Expand source code</span>
|
||
</summary>
|
||
<pre><code class="python">@dataclass
|
||
class Sample:
|
||
id: str
|
||
path: str
|
||
epoch: int
|
||
duration: float
|
||
conditioning: tp.Optional[tp.Dict[str, tp.Any]]
|
||
prompt: tp.Optional[ReferenceSample]
|
||
reference: tp.Optional[ReferenceSample]
|
||
generation_args: tp.Optional[tp.Dict[str, tp.Any]]
|
||
|
||
def __hash__(self):
|
||
return hash(self.id)
|
||
|
||
def audio(self) -> tp.Tuple[torch.Tensor, int]:
|
||
return audio_read(self.path)
|
||
|
||
def audio_prompt(self) -> tp.Optional[tp.Tuple[torch.Tensor, int]]:
|
||
return audio_read(self.prompt.path) if self.prompt is not None else None
|
||
|
||
def audio_reference(self) -> tp.Optional[tp.Tuple[torch.Tensor, int]]:
|
||
return audio_read(self.reference.path) if self.reference is not None else None</code></pre>
|
||
</details>
|
||
<div class="desc"><p>Sample(id: str, path: str, epoch: int, duration: float, conditioning: Optional[Dict[str, Any]], prompt: Optional[audiocraft.utils.samples.manager.ReferenceSample], reference: Optional[audiocraft.utils.samples.manager.ReferenceSample], generation_args: Optional[Dict[str, Any]])</p></div>
|
||
<h3>Class variables</h3>
|
||
<dl>
|
||
<dt id="audiocraft.utils.samples.manager.Sample.conditioning"><code class="name">var <span class="ident">conditioning</span> : Dict[str, Any] | None</code></dt>
|
||
<dd>
|
||
<div class="desc"></div>
|
||
</dd>
|
||
<dt id="audiocraft.utils.samples.manager.Sample.duration"><code class="name">var <span class="ident">duration</span> : float</code></dt>
|
||
<dd>
|
||
<div class="desc"></div>
|
||
</dd>
|
||
<dt id="audiocraft.utils.samples.manager.Sample.epoch"><code class="name">var <span class="ident">epoch</span> : int</code></dt>
|
||
<dd>
|
||
<div class="desc"></div>
|
||
</dd>
|
||
<dt id="audiocraft.utils.samples.manager.Sample.generation_args"><code class="name">var <span class="ident">generation_args</span> : Dict[str, Any] | None</code></dt>
|
||
<dd>
|
||
<div class="desc"></div>
|
||
</dd>
|
||
<dt id="audiocraft.utils.samples.manager.Sample.id"><code class="name">var <span class="ident">id</span> : str</code></dt>
|
||
<dd>
|
||
<div class="desc"></div>
|
||
</dd>
|
||
<dt id="audiocraft.utils.samples.manager.Sample.path"><code class="name">var <span class="ident">path</span> : str</code></dt>
|
||
<dd>
|
||
<div class="desc"></div>
|
||
</dd>
|
||
<dt id="audiocraft.utils.samples.manager.Sample.prompt"><code class="name">var <span class="ident">prompt</span> : <a title="audiocraft.utils.samples.manager.ReferenceSample" href="#audiocraft.utils.samples.manager.ReferenceSample">ReferenceSample</a> | None</code></dt>
|
||
<dd>
|
||
<div class="desc"></div>
|
||
</dd>
|
||
<dt id="audiocraft.utils.samples.manager.Sample.reference"><code class="name">var <span class="ident">reference</span> : <a title="audiocraft.utils.samples.manager.ReferenceSample" href="#audiocraft.utils.samples.manager.ReferenceSample">ReferenceSample</a> | None</code></dt>
|
||
<dd>
|
||
<div class="desc"></div>
|
||
</dd>
|
||
</dl>
|
||
<h3>Methods</h3>
|
||
<dl>
|
||
<dt id="audiocraft.utils.samples.manager.Sample.audio"><code class="name flex">
|
||
<span>def <span class="ident">audio</span></span>(<span>self) ‑> Tuple[torch.Tensor, int]</span>
|
||
</code></dt>
|
||
<dd>
|
||
<details class="source">
|
||
<summary>
|
||
<span>Expand source code</span>
|
||
</summary>
|
||
<pre><code class="python">def audio(self) -> tp.Tuple[torch.Tensor, int]:
|
||
return audio_read(self.path)</code></pre>
|
||
</details>
|
||
<div class="desc"></div>
|
||
</dd>
|
||
<dt id="audiocraft.utils.samples.manager.Sample.audio_prompt"><code class="name flex">
|
||
<span>def <span class="ident">audio_prompt</span></span>(<span>self) ‑> Tuple[torch.Tensor, int] | None</span>
|
||
</code></dt>
|
||
<dd>
|
||
<details class="source">
|
||
<summary>
|
||
<span>Expand source code</span>
|
||
</summary>
|
||
<pre><code class="python">def audio_prompt(self) -> tp.Optional[tp.Tuple[torch.Tensor, int]]:
|
||
return audio_read(self.prompt.path) if self.prompt is not None else None</code></pre>
|
||
</details>
|
||
<div class="desc"></div>
|
||
</dd>
|
||
<dt id="audiocraft.utils.samples.manager.Sample.audio_reference"><code class="name flex">
|
||
<span>def <span class="ident">audio_reference</span></span>(<span>self) ‑> Tuple[torch.Tensor, int] | None</span>
|
||
</code></dt>
|
||
<dd>
|
||
<details class="source">
|
||
<summary>
|
||
<span>Expand source code</span>
|
||
</summary>
|
||
<pre><code class="python">def audio_reference(self) -> tp.Optional[tp.Tuple[torch.Tensor, int]]:
|
||
return audio_read(self.reference.path) if self.reference is not None else None</code></pre>
|
||
</details>
|
||
<div class="desc"></div>
|
||
</dd>
|
||
</dl>
|
||
</dd>
|
||
<dt id="audiocraft.utils.samples.manager.SampleManager"><code class="flex name class">
|
||
<span>class <span class="ident">SampleManager</span></span>
|
||
<span>(</span><span>xp: dora.xp.XP, map_reference_to_sample_id: bool = False)</span>
|
||
</code></dt>
|
||
<dd>
|
||
<details class="source">
|
||
<summary>
|
||
<span>Expand source code</span>
|
||
</summary>
|
||
<pre><code class="python">class SampleManager:
|
||
"""Audio samples IO handling within a given dora xp.
|
||
|
||
The sample manager handles the dumping and loading logic for generated and
|
||
references samples across epochs for a given xp, providing a simple API to
|
||
store, retrieve and compare audio samples.
|
||
|
||
Args:
|
||
xp (dora.XP): Dora experiment object. The XP contains information on the XP folder
|
||
where all outputs are stored and the configuration of the experiment,
|
||
which is useful to retrieve audio-related parameters.
|
||
map_reference_to_sample_id (bool): Whether to use the sample_id for all reference samples
|
||
instead of generating a dedicated hash id. This is useful to allow easier comparison
|
||
with ground truth sample from the files directly without having to read the JSON metadata
|
||
to do the mapping (at the cost of potentially dumping duplicate prompts/references
|
||
depending on the task).
|
||
"""
|
||
def __init__(self, xp: dora.XP, map_reference_to_sample_id: bool = False):
|
||
self.xp = xp
|
||
self.base_folder: Path = xp.folder / xp.cfg.generate.path
|
||
self.reference_folder = self.base_folder / 'reference'
|
||
self.map_reference_to_sample_id = map_reference_to_sample_id
|
||
self.samples: tp.List[Sample] = []
|
||
self._load_samples()
|
||
|
||
@property
|
||
def latest_epoch(self):
|
||
"""Latest epoch across all samples."""
|
||
return max(self.samples, key=lambda x: x.epoch).epoch if self.samples else 0
|
||
|
||
def _load_samples(self):
|
||
"""Scan the sample folder and load existing samples."""
|
||
jsons = self.base_folder.glob('**/*.json')
|
||
with ThreadPoolExecutor(6) as pool:
|
||
self.samples = list(pool.map(self._load_sample, jsons))
|
||
|
||
@staticmethod
|
||
@lru_cache(2**26)
|
||
def _load_sample(json_file: Path) -> Sample:
|
||
with open(json_file, 'r') as f:
|
||
data: tp.Dict[str, tp.Any] = json.load(f)
|
||
# fetch prompt data
|
||
prompt_data = data.get('prompt')
|
||
prompt = ReferenceSample(id=prompt_data['id'], path=prompt_data['path'],
|
||
duration=prompt_data['duration']) if prompt_data else None
|
||
# fetch reference data
|
||
reference_data = data.get('reference')
|
||
reference = ReferenceSample(id=reference_data['id'], path=reference_data['path'],
|
||
duration=reference_data['duration']) if reference_data else None
|
||
# build sample object
|
||
return Sample(id=data['id'], path=data['path'], epoch=data['epoch'], duration=data['duration'],
|
||
prompt=prompt, conditioning=data.get('conditioning'), reference=reference,
|
||
generation_args=data.get('generation_args'))
|
||
|
||
def _init_hash(self):
|
||
return hashlib.sha1()
|
||
|
||
def _get_tensor_id(self, tensor: torch.Tensor) -> str:
|
||
hash_id = self._init_hash()
|
||
hash_id.update(tensor.numpy().data)
|
||
return hash_id.hexdigest()
|
||
|
||
def _get_sample_id(self, index: int, prompt_wav: tp.Optional[torch.Tensor],
|
||
conditions: tp.Optional[tp.Dict[str, str]]) -> str:
|
||
"""Computes an id for a sample given its input data.
|
||
This id is deterministic if prompt and/or conditions are provided by using a sha1 hash on the input.
|
||
Otherwise, a random id of the form "noinput_{uuid4().hex}" is returned.
|
||
|
||
Args:
|
||
index (int): Batch index, Helpful to differentiate samples from the same batch.
|
||
prompt_wav (torch.Tensor): Prompt used during generation.
|
||
conditions (dict[str, str]): Conditioning used during generation.
|
||
"""
|
||
# For totally unconditioned generations we will just use a random UUID.
|
||
# The function get_samples_for_xps will do a simple ordered match with a custom key.
|
||
if prompt_wav is None and not conditions:
|
||
return f"noinput_{uuid.uuid4().hex}"
|
||
|
||
# Human readable portion
|
||
hr_label = ""
|
||
# Create a deterministic id using hashing
|
||
hash_id = self._init_hash()
|
||
hash_id.update(f"{index}".encode())
|
||
if prompt_wav is not None:
|
||
hash_id.update(prompt_wav.numpy().data)
|
||
hr_label += "_prompted"
|
||
else:
|
||
hr_label += "_unprompted"
|
||
if conditions:
|
||
encoded_json = json.dumps(conditions, sort_keys=True).encode()
|
||
hash_id.update(encoded_json)
|
||
cond_str = "-".join([f"{key}={slugify(value)}"
|
||
for key, value in sorted(conditions.items())])
|
||
cond_str = cond_str[:100] # some raw text might be too long to be a valid filename
|
||
cond_str = cond_str if len(cond_str) > 0 else "unconditioned"
|
||
hr_label += f"_{cond_str}"
|
||
else:
|
||
hr_label += "_unconditioned"
|
||
|
||
return hash_id.hexdigest() + hr_label
|
||
|
||
def _store_audio(self, wav: torch.Tensor, stem_path: Path, overwrite: bool = False) -> Path:
|
||
"""Stores the audio with the given stem path using the XP's configuration.
|
||
|
||
Args:
|
||
wav (torch.Tensor): Audio to store.
|
||
stem_path (Path): Path in sample output directory with file stem to use.
|
||
overwrite (bool): When False (default), skips storing an existing audio file.
|
||
Returns:
|
||
Path: The path at which the audio is stored.
|
||
"""
|
||
existing_paths = [
|
||
path for path in stem_path.parent.glob(stem_path.stem + '.*')
|
||
if path.suffix != '.json'
|
||
]
|
||
exists = len(existing_paths) > 0
|
||
if exists and overwrite:
|
||
logger.warning(f"Overwriting existing audio file with stem path {stem_path}")
|
||
elif exists:
|
||
return existing_paths[0]
|
||
|
||
audio_path = audio_write(stem_path, wav, **self.xp.cfg.generate.audio)
|
||
return audio_path
|
||
|
||
def add_sample(self, sample_wav: torch.Tensor, epoch: int, index: int = 0,
|
||
conditions: tp.Optional[tp.Dict[str, str]] = None, prompt_wav: tp.Optional[torch.Tensor] = None,
|
||
ground_truth_wav: tp.Optional[torch.Tensor] = None,
|
||
generation_args: tp.Optional[tp.Dict[str, tp.Any]] = None) -> Sample:
|
||
"""Adds a single sample.
|
||
The sample is stored in the XP's sample output directory, under a corresponding epoch folder.
|
||
Each sample is assigned an id which is computed using the input data. In addition to the
|
||
sample itself, a json file containing associated metadata is stored next to it.
|
||
|
||
Args:
|
||
sample_wav (torch.Tensor): sample audio to store. Tensor of shape [channels, shape].
|
||
epoch (int): current training epoch.
|
||
index (int): helpful to differentiate samples from the same batch.
|
||
conditions (dict[str, str], optional): conditioning used during generation.
|
||
prompt_wav (torch.Tensor, optional): prompt used during generation. Tensor of shape [channels, shape].
|
||
ground_truth_wav (torch.Tensor, optional): reference audio where prompt was extracted from.
|
||
Tensor of shape [channels, shape].
|
||
generation_args (dict[str, any], optional): dictionary of other arguments used during generation.
|
||
Returns:
|
||
Sample: The saved sample.
|
||
"""
|
||
sample_id = self._get_sample_id(index, prompt_wav, conditions)
|
||
reuse_id = self.map_reference_to_sample_id
|
||
prompt, ground_truth = None, None
|
||
if prompt_wav is not None:
|
||
prompt_id = sample_id if reuse_id else self._get_tensor_id(prompt_wav.sum(0, keepdim=True))
|
||
prompt_duration = prompt_wav.shape[-1] / self.xp.cfg.sample_rate
|
||
prompt_path = self._store_audio(prompt_wav, self.base_folder / str(epoch) / 'prompt' / prompt_id)
|
||
prompt = ReferenceSample(prompt_id, str(prompt_path), prompt_duration)
|
||
if ground_truth_wav is not None:
|
||
ground_truth_id = sample_id if reuse_id else self._get_tensor_id(ground_truth_wav.sum(0, keepdim=True))
|
||
ground_truth_duration = ground_truth_wav.shape[-1] / self.xp.cfg.sample_rate
|
||
ground_truth_path = self._store_audio(ground_truth_wav, self.base_folder / 'reference' / ground_truth_id)
|
||
ground_truth = ReferenceSample(ground_truth_id, str(ground_truth_path), ground_truth_duration)
|
||
sample_path = self._store_audio(sample_wav, self.base_folder / str(epoch) / sample_id, overwrite=True)
|
||
duration = sample_wav.shape[-1] / self.xp.cfg.sample_rate
|
||
sample = Sample(sample_id, str(sample_path), epoch, duration, conditions, prompt, ground_truth, generation_args)
|
||
self.samples.append(sample)
|
||
with open(sample_path.with_suffix('.json'), 'w') as f:
|
||
json.dump(asdict(sample), f, indent=2)
|
||
return sample
|
||
|
||
def add_samples(self, samples_wavs: torch.Tensor, epoch: int,
|
||
conditioning: tp.Optional[tp.List[tp.Dict[str, tp.Any]]] = None,
|
||
prompt_wavs: tp.Optional[torch.Tensor] = None,
|
||
ground_truth_wavs: tp.Optional[torch.Tensor] = None,
|
||
generation_args: tp.Optional[tp.Dict[str, tp.Any]] = None) -> tp.List[Sample]:
|
||
"""Adds a batch of samples.
|
||
The samples are stored in the XP's sample output directory, under a corresponding
|
||
epoch folder. Each sample is assigned an id which is computed using the input data and their batch index.
|
||
In addition to the sample itself, a json file containing associated metadata is stored next to it.
|
||
|
||
Args:
|
||
sample_wavs (torch.Tensor): Batch of audio wavs to store. Tensor of shape [batch_size, channels, shape].
|
||
epoch (int): Current training epoch.
|
||
conditioning (list of dict[str, str], optional): List of conditions used during generation,
|
||
one per sample in the batch.
|
||
prompt_wavs (torch.Tensor, optional): Prompts used during generation. Tensor of shape
|
||
[batch_size, channels, shape].
|
||
ground_truth_wav (torch.Tensor, optional): Reference audio where prompts were extracted from.
|
||
Tensor of shape [batch_size, channels, shape].
|
||
generation_args (dict[str, Any], optional): Dictionary of other arguments used during generation.
|
||
Returns:
|
||
samples (list of Sample): The saved audio samples with prompts, ground truth and metadata.
|
||
"""
|
||
samples = []
|
||
for idx, wav in enumerate(samples_wavs):
|
||
prompt_wav = prompt_wavs[idx] if prompt_wavs is not None else None
|
||
gt_wav = ground_truth_wavs[idx] if ground_truth_wavs is not None else None
|
||
conditions = conditioning[idx] if conditioning is not None else None
|
||
samples.append(self.add_sample(wav, epoch, idx, conditions, prompt_wav, gt_wav, generation_args))
|
||
return samples
|
||
|
||
def get_samples(self, epoch: int = -1, max_epoch: int = -1, exclude_prompted: bool = False,
|
||
exclude_unprompted: bool = False, exclude_conditioned: bool = False,
|
||
exclude_unconditioned: bool = False) -> tp.Set[Sample]:
|
||
"""Returns a set of samples for this XP. Optionally, you can filter which samples to obtain.
|
||
Please note that existing samples are loaded during the manager's initialization, and added samples through this
|
||
manager are also tracked. Any other external changes are not tracked automatically, so creating a new manager
|
||
is the only way detect them.
|
||
|
||
Args:
|
||
epoch (int): If provided, only return samples corresponding to this epoch.
|
||
max_epoch (int): If provided, only return samples corresponding to the latest epoch that is <= max_epoch.
|
||
exclude_prompted (bool): If True, does not include samples that used a prompt.
|
||
exclude_unprompted (bool): If True, does not include samples that did not use a prompt.
|
||
exclude_conditioned (bool): If True, excludes samples that used conditioning.
|
||
exclude_unconditioned (bool): If True, excludes samples that did not use conditioning.
|
||
Returns:
|
||
Samples (set of Sample): The retrieved samples matching the provided filters.
|
||
"""
|
||
if max_epoch >= 0:
|
||
samples_epoch = max(sample.epoch for sample in self.samples if sample.epoch <= max_epoch)
|
||
else:
|
||
samples_epoch = self.latest_epoch if epoch < 0 else epoch
|
||
samples = {
|
||
sample
|
||
for sample in self.samples
|
||
if (
|
||
(sample.epoch == samples_epoch) and
|
||
(not exclude_prompted or sample.prompt is None) and
|
||
(not exclude_unprompted or sample.prompt is not None) and
|
||
(not exclude_conditioned or not sample.conditioning) and
|
||
(not exclude_unconditioned or sample.conditioning)
|
||
)
|
||
}
|
||
return samples</code></pre>
|
||
</details>
|
||
<div class="desc"><p>Audio samples IO handling within a given dora xp.</p>
|
||
<p>The sample manager handles the dumping and loading logic for generated and
|
||
references samples across epochs for a given xp, providing a simple API to
|
||
store, retrieve and compare audio samples.</p>
|
||
<h2 id="args">Args</h2>
|
||
<dl>
|
||
<dt><strong><code>xp</code></strong> : <code>dora.XP</code></dt>
|
||
<dd>Dora experiment object. The XP contains information on the XP folder
|
||
where all outputs are stored and the configuration of the experiment,
|
||
which is useful to retrieve audio-related parameters.</dd>
|
||
<dt><strong><code>map_reference_to_sample_id</code></strong> : <code>bool</code></dt>
|
||
<dd>Whether to use the sample_id for all reference samples
|
||
instead of generating a dedicated hash id. This is useful to allow easier comparison
|
||
with ground truth sample from the files directly without having to read the JSON metadata
|
||
to do the mapping (at the cost of potentially dumping duplicate prompts/references
|
||
depending on the task).</dd>
|
||
</dl></div>
|
||
<h3>Instance variables</h3>
|
||
<dl>
|
||
<dt id="audiocraft.utils.samples.manager.SampleManager.latest_epoch"><code class="name">prop <span class="ident">latest_epoch</span></code></dt>
|
||
<dd>
|
||
<details class="source">
|
||
<summary>
|
||
<span>Expand source code</span>
|
||
</summary>
|
||
<pre><code class="python">@property
|
||
def latest_epoch(self):
|
||
"""Latest epoch across all samples."""
|
||
return max(self.samples, key=lambda x: x.epoch).epoch if self.samples else 0</code></pre>
|
||
</details>
|
||
<div class="desc"><p>Latest epoch across all samples.</p></div>
|
||
</dd>
|
||
</dl>
|
||
<h3>Methods</h3>
|
||
<dl>
|
||
<dt id="audiocraft.utils.samples.manager.SampleManager.add_sample"><code class="name flex">
|
||
<span>def <span class="ident">add_sample</span></span>(<span>self,<br>sample_wav: torch.Tensor,<br>epoch: int,<br>index: int = 0,<br>conditions: Dict[str, str] | None = None,<br>prompt_wav: torch.Tensor | None = None,<br>ground_truth_wav: torch.Tensor | None = None,<br>generation_args: Dict[str, Any] | None = None) ‑> <a title="audiocraft.utils.samples.manager.Sample" href="#audiocraft.utils.samples.manager.Sample">Sample</a></span>
|
||
</code></dt>
|
||
<dd>
|
||
<details class="source">
|
||
<summary>
|
||
<span>Expand source code</span>
|
||
</summary>
|
||
<pre><code class="python">def add_sample(self, sample_wav: torch.Tensor, epoch: int, index: int = 0,
|
||
conditions: tp.Optional[tp.Dict[str, str]] = None, prompt_wav: tp.Optional[torch.Tensor] = None,
|
||
ground_truth_wav: tp.Optional[torch.Tensor] = None,
|
||
generation_args: tp.Optional[tp.Dict[str, tp.Any]] = None) -> Sample:
|
||
"""Adds a single sample.
|
||
The sample is stored in the XP's sample output directory, under a corresponding epoch folder.
|
||
Each sample is assigned an id which is computed using the input data. In addition to the
|
||
sample itself, a json file containing associated metadata is stored next to it.
|
||
|
||
Args:
|
||
sample_wav (torch.Tensor): sample audio to store. Tensor of shape [channels, shape].
|
||
epoch (int): current training epoch.
|
||
index (int): helpful to differentiate samples from the same batch.
|
||
conditions (dict[str, str], optional): conditioning used during generation.
|
||
prompt_wav (torch.Tensor, optional): prompt used during generation. Tensor of shape [channels, shape].
|
||
ground_truth_wav (torch.Tensor, optional): reference audio where prompt was extracted from.
|
||
Tensor of shape [channels, shape].
|
||
generation_args (dict[str, any], optional): dictionary of other arguments used during generation.
|
||
Returns:
|
||
Sample: The saved sample.
|
||
"""
|
||
sample_id = self._get_sample_id(index, prompt_wav, conditions)
|
||
reuse_id = self.map_reference_to_sample_id
|
||
prompt, ground_truth = None, None
|
||
if prompt_wav is not None:
|
||
prompt_id = sample_id if reuse_id else self._get_tensor_id(prompt_wav.sum(0, keepdim=True))
|
||
prompt_duration = prompt_wav.shape[-1] / self.xp.cfg.sample_rate
|
||
prompt_path = self._store_audio(prompt_wav, self.base_folder / str(epoch) / 'prompt' / prompt_id)
|
||
prompt = ReferenceSample(prompt_id, str(prompt_path), prompt_duration)
|
||
if ground_truth_wav is not None:
|
||
ground_truth_id = sample_id if reuse_id else self._get_tensor_id(ground_truth_wav.sum(0, keepdim=True))
|
||
ground_truth_duration = ground_truth_wav.shape[-1] / self.xp.cfg.sample_rate
|
||
ground_truth_path = self._store_audio(ground_truth_wav, self.base_folder / 'reference' / ground_truth_id)
|
||
ground_truth = ReferenceSample(ground_truth_id, str(ground_truth_path), ground_truth_duration)
|
||
sample_path = self._store_audio(sample_wav, self.base_folder / str(epoch) / sample_id, overwrite=True)
|
||
duration = sample_wav.shape[-1] / self.xp.cfg.sample_rate
|
||
sample = Sample(sample_id, str(sample_path), epoch, duration, conditions, prompt, ground_truth, generation_args)
|
||
self.samples.append(sample)
|
||
with open(sample_path.with_suffix('.json'), 'w') as f:
|
||
json.dump(asdict(sample), f, indent=2)
|
||
return sample</code></pre>
|
||
</details>
|
||
<div class="desc"><p>Adds a single sample.
|
||
The sample is stored in the XP's sample output directory, under a corresponding epoch folder.
|
||
Each sample is assigned an id which is computed using the input data. In addition to the
|
||
sample itself, a json file containing associated metadata is stored next to it.</p>
|
||
<h2 id="args">Args</h2>
|
||
<dl>
|
||
<dt><strong><code>sample_wav</code></strong> : <code>torch.Tensor</code></dt>
|
||
<dd>sample audio to store. Tensor of shape [channels, shape].</dd>
|
||
<dt><strong><code>epoch</code></strong> : <code>int</code></dt>
|
||
<dd>current training epoch.</dd>
|
||
<dt><strong><code>index</code></strong> : <code>int</code></dt>
|
||
<dd>helpful to differentiate samples from the same batch.</dd>
|
||
<dt><strong><code>conditions</code></strong> : <code>dict[str, str]</code>, optional</dt>
|
||
<dd>conditioning used during generation.</dd>
|
||
<dt><strong><code>prompt_wav</code></strong> : <code>torch.Tensor</code>, optional</dt>
|
||
<dd>prompt used during generation. Tensor of shape [channels, shape].</dd>
|
||
<dt><strong><code>ground_truth_wav</code></strong> : <code>torch.Tensor</code>, optional</dt>
|
||
<dd>reference audio where prompt was extracted from.
|
||
Tensor of shape [channels, shape].</dd>
|
||
<dt><strong><code>generation_args</code></strong> : <code>dict[str, any]</code>, optional</dt>
|
||
<dd>dictionary of other arguments used during generation.</dd>
|
||
</dl>
|
||
<h2 id="returns">Returns</h2>
|
||
<dl>
|
||
<dt><code><a title="audiocraft.utils.samples.manager.Sample" href="#audiocraft.utils.samples.manager.Sample">Sample</a></code></dt>
|
||
<dd>The saved sample.</dd>
|
||
</dl></div>
|
||
</dd>
|
||
<dt id="audiocraft.utils.samples.manager.SampleManager.add_samples"><code class="name flex">
|
||
<span>def <span class="ident">add_samples</span></span>(<span>self,<br>samples_wavs: torch.Tensor,<br>epoch: int,<br>conditioning: List[Dict[str, Any]] | None = None,<br>prompt_wavs: torch.Tensor | None = None,<br>ground_truth_wavs: torch.Tensor | None = None,<br>generation_args: Dict[str, Any] | None = None) ‑> List[<a title="audiocraft.utils.samples.manager.Sample" href="#audiocraft.utils.samples.manager.Sample">Sample</a>]</span>
|
||
</code></dt>
|
||
<dd>
|
||
<details class="source">
|
||
<summary>
|
||
<span>Expand source code</span>
|
||
</summary>
|
||
<pre><code class="python">def add_samples(self, samples_wavs: torch.Tensor, epoch: int,
|
||
conditioning: tp.Optional[tp.List[tp.Dict[str, tp.Any]]] = None,
|
||
prompt_wavs: tp.Optional[torch.Tensor] = None,
|
||
ground_truth_wavs: tp.Optional[torch.Tensor] = None,
|
||
generation_args: tp.Optional[tp.Dict[str, tp.Any]] = None) -> tp.List[Sample]:
|
||
"""Adds a batch of samples.
|
||
The samples are stored in the XP's sample output directory, under a corresponding
|
||
epoch folder. Each sample is assigned an id which is computed using the input data and their batch index.
|
||
In addition to the sample itself, a json file containing associated metadata is stored next to it.
|
||
|
||
Args:
|
||
sample_wavs (torch.Tensor): Batch of audio wavs to store. Tensor of shape [batch_size, channels, shape].
|
||
epoch (int): Current training epoch.
|
||
conditioning (list of dict[str, str], optional): List of conditions used during generation,
|
||
one per sample in the batch.
|
||
prompt_wavs (torch.Tensor, optional): Prompts used during generation. Tensor of shape
|
||
[batch_size, channels, shape].
|
||
ground_truth_wav (torch.Tensor, optional): Reference audio where prompts were extracted from.
|
||
Tensor of shape [batch_size, channels, shape].
|
||
generation_args (dict[str, Any], optional): Dictionary of other arguments used during generation.
|
||
Returns:
|
||
samples (list of Sample): The saved audio samples with prompts, ground truth and metadata.
|
||
"""
|
||
samples = []
|
||
for idx, wav in enumerate(samples_wavs):
|
||
prompt_wav = prompt_wavs[idx] if prompt_wavs is not None else None
|
||
gt_wav = ground_truth_wavs[idx] if ground_truth_wavs is not None else None
|
||
conditions = conditioning[idx] if conditioning is not None else None
|
||
samples.append(self.add_sample(wav, epoch, idx, conditions, prompt_wav, gt_wav, generation_args))
|
||
return samples</code></pre>
|
||
</details>
|
||
<div class="desc"><p>Adds a batch of samples.
|
||
The samples are stored in the XP's sample output directory, under a corresponding
|
||
epoch folder. Each sample is assigned an id which is computed using the input data and their batch index.
|
||
In addition to the sample itself, a json file containing associated metadata is stored next to it.</p>
|
||
<h2 id="args">Args</h2>
|
||
<dl>
|
||
<dt><strong><code>sample_wavs</code></strong> : <code>torch.Tensor</code></dt>
|
||
<dd>Batch of audio wavs to store. Tensor of shape [batch_size, channels, shape].</dd>
|
||
<dt><strong><code>epoch</code></strong> : <code>int</code></dt>
|
||
<dd>Current training epoch.</dd>
|
||
<dt><strong><code>conditioning</code></strong> : <code>list</code> of <code>dict[str, str]</code>, optional</dt>
|
||
<dd>List of conditions used during generation,
|
||
one per sample in the batch.</dd>
|
||
<dt><strong><code>prompt_wavs</code></strong> : <code>torch.Tensor</code>, optional</dt>
|
||
<dd>Prompts used during generation. Tensor of shape
|
||
[batch_size, channels, shape].</dd>
|
||
<dt><strong><code>ground_truth_wav</code></strong> : <code>torch.Tensor</code>, optional</dt>
|
||
<dd>Reference audio where prompts were extracted from.
|
||
Tensor of shape [batch_size, channels, shape].</dd>
|
||
<dt><strong><code>generation_args</code></strong> : <code>dict[str, Any]</code>, optional</dt>
|
||
<dd>Dictionary of other arguments used during generation.</dd>
|
||
</dl>
|
||
<h2 id="returns">Returns</h2>
|
||
<p>samples (list of Sample): The saved audio samples with prompts, ground truth and metadata.</p></div>
|
||
</dd>
|
||
<dt id="audiocraft.utils.samples.manager.SampleManager.get_samples"><code class="name flex">
|
||
<span>def <span class="ident">get_samples</span></span>(<span>self,<br>epoch: int = -1,<br>max_epoch: int = -1,<br>exclude_prompted: bool = False,<br>exclude_unprompted: bool = False,<br>exclude_conditioned: bool = False,<br>exclude_unconditioned: bool = False) ‑> Set[<a title="audiocraft.utils.samples.manager.Sample" href="#audiocraft.utils.samples.manager.Sample">Sample</a>]</span>
|
||
</code></dt>
|
||
<dd>
|
||
<details class="source">
|
||
<summary>
|
||
<span>Expand source code</span>
|
||
</summary>
|
||
<pre><code class="python">def get_samples(self, epoch: int = -1, max_epoch: int = -1, exclude_prompted: bool = False,
|
||
exclude_unprompted: bool = False, exclude_conditioned: bool = False,
|
||
exclude_unconditioned: bool = False) -> tp.Set[Sample]:
|
||
"""Returns a set of samples for this XP. Optionally, you can filter which samples to obtain.
|
||
Please note that existing samples are loaded during the manager's initialization, and added samples through this
|
||
manager are also tracked. Any other external changes are not tracked automatically, so creating a new manager
|
||
is the only way detect them.
|
||
|
||
Args:
|
||
epoch (int): If provided, only return samples corresponding to this epoch.
|
||
max_epoch (int): If provided, only return samples corresponding to the latest epoch that is <= max_epoch.
|
||
exclude_prompted (bool): If True, does not include samples that used a prompt.
|
||
exclude_unprompted (bool): If True, does not include samples that did not use a prompt.
|
||
exclude_conditioned (bool): If True, excludes samples that used conditioning.
|
||
exclude_unconditioned (bool): If True, excludes samples that did not use conditioning.
|
||
Returns:
|
||
Samples (set of Sample): The retrieved samples matching the provided filters.
|
||
"""
|
||
if max_epoch >= 0:
|
||
samples_epoch = max(sample.epoch for sample in self.samples if sample.epoch <= max_epoch)
|
||
else:
|
||
samples_epoch = self.latest_epoch if epoch < 0 else epoch
|
||
samples = {
|
||
sample
|
||
for sample in self.samples
|
||
if (
|
||
(sample.epoch == samples_epoch) and
|
||
(not exclude_prompted or sample.prompt is None) and
|
||
(not exclude_unprompted or sample.prompt is not None) and
|
||
(not exclude_conditioned or not sample.conditioning) and
|
||
(not exclude_unconditioned or sample.conditioning)
|
||
)
|
||
}
|
||
return samples</code></pre>
|
||
</details>
|
||
<div class="desc"><p>Returns a set of samples for this XP. Optionally, you can filter which samples to obtain.
|
||
Please note that existing samples are loaded during the manager's initialization, and added samples through this
|
||
manager are also tracked. Any other external changes are not tracked automatically, so creating a new manager
|
||
is the only way detect them.</p>
|
||
<h2 id="args">Args</h2>
|
||
<dl>
|
||
<dt><strong><code>epoch</code></strong> : <code>int</code></dt>
|
||
<dd>If provided, only return samples corresponding to this epoch.</dd>
|
||
<dt><strong><code>max_epoch</code></strong> : <code>int</code></dt>
|
||
<dd>If provided, only return samples corresponding to the latest epoch that is <= max_epoch.</dd>
|
||
<dt><strong><code>exclude_prompted</code></strong> : <code>bool</code></dt>
|
||
<dd>If True, does not include samples that used a prompt.</dd>
|
||
<dt><strong><code>exclude_unprompted</code></strong> : <code>bool</code></dt>
|
||
<dd>If True, does not include samples that did not use a prompt.</dd>
|
||
<dt><strong><code>exclude_conditioned</code></strong> : <code>bool</code></dt>
|
||
<dd>If True, excludes samples that used conditioning.</dd>
|
||
<dt><strong><code>exclude_unconditioned</code></strong> : <code>bool</code></dt>
|
||
<dd>If True, excludes samples that did not use conditioning.</dd>
|
||
</dl>
|
||
<h2 id="returns">Returns</h2>
|
||
<p>Samples (set of Sample): The retrieved samples matching the provided filters.</p></div>
|
||
</dd>
|
||
</dl>
|
||
</dd>
|
||
</dl>
|
||
</section>
|
||
</article>
|
||
<nav id="sidebar">
|
||
<div class="toc">
|
||
<ul></ul>
|
||
</div>
|
||
<ul id="index">
|
||
<li><h3>Super-module</h3>
|
||
<ul>
|
||
<li><code><a title="audiocraft.utils.samples" href="index.html">audiocraft.utils.samples</a></code></li>
|
||
</ul>
|
||
</li>
|
||
<li><h3><a href="#header-functions">Functions</a></h3>
|
||
<ul class="">
|
||
<li><code><a title="audiocraft.utils.samples.manager.get_samples_for_xps" href="#audiocraft.utils.samples.manager.get_samples_for_xps">get_samples_for_xps</a></code></li>
|
||
<li><code><a title="audiocraft.utils.samples.manager.slugify" href="#audiocraft.utils.samples.manager.slugify">slugify</a></code></li>
|
||
</ul>
|
||
</li>
|
||
<li><h3><a href="#header-classes">Classes</a></h3>
|
||
<ul>
|
||
<li>
|
||
<h4><code><a title="audiocraft.utils.samples.manager.ReferenceSample" href="#audiocraft.utils.samples.manager.ReferenceSample">ReferenceSample</a></code></h4>
|
||
<ul class="">
|
||
<li><code><a title="audiocraft.utils.samples.manager.ReferenceSample.duration" href="#audiocraft.utils.samples.manager.ReferenceSample.duration">duration</a></code></li>
|
||
<li><code><a title="audiocraft.utils.samples.manager.ReferenceSample.id" href="#audiocraft.utils.samples.manager.ReferenceSample.id">id</a></code></li>
|
||
<li><code><a title="audiocraft.utils.samples.manager.ReferenceSample.path" href="#audiocraft.utils.samples.manager.ReferenceSample.path">path</a></code></li>
|
||
</ul>
|
||
</li>
|
||
<li>
|
||
<h4><code><a title="audiocraft.utils.samples.manager.Sample" href="#audiocraft.utils.samples.manager.Sample">Sample</a></code></h4>
|
||
<ul class="two-column">
|
||
<li><code><a title="audiocraft.utils.samples.manager.Sample.audio" href="#audiocraft.utils.samples.manager.Sample.audio">audio</a></code></li>
|
||
<li><code><a title="audiocraft.utils.samples.manager.Sample.audio_prompt" href="#audiocraft.utils.samples.manager.Sample.audio_prompt">audio_prompt</a></code></li>
|
||
<li><code><a title="audiocraft.utils.samples.manager.Sample.audio_reference" href="#audiocraft.utils.samples.manager.Sample.audio_reference">audio_reference</a></code></li>
|
||
<li><code><a title="audiocraft.utils.samples.manager.Sample.conditioning" href="#audiocraft.utils.samples.manager.Sample.conditioning">conditioning</a></code></li>
|
||
<li><code><a title="audiocraft.utils.samples.manager.Sample.duration" href="#audiocraft.utils.samples.manager.Sample.duration">duration</a></code></li>
|
||
<li><code><a title="audiocraft.utils.samples.manager.Sample.epoch" href="#audiocraft.utils.samples.manager.Sample.epoch">epoch</a></code></li>
|
||
<li><code><a title="audiocraft.utils.samples.manager.Sample.generation_args" href="#audiocraft.utils.samples.manager.Sample.generation_args">generation_args</a></code></li>
|
||
<li><code><a title="audiocraft.utils.samples.manager.Sample.id" href="#audiocraft.utils.samples.manager.Sample.id">id</a></code></li>
|
||
<li><code><a title="audiocraft.utils.samples.manager.Sample.path" href="#audiocraft.utils.samples.manager.Sample.path">path</a></code></li>
|
||
<li><code><a title="audiocraft.utils.samples.manager.Sample.prompt" href="#audiocraft.utils.samples.manager.Sample.prompt">prompt</a></code></li>
|
||
<li><code><a title="audiocraft.utils.samples.manager.Sample.reference" href="#audiocraft.utils.samples.manager.Sample.reference">reference</a></code></li>
|
||
</ul>
|
||
</li>
|
||
<li>
|
||
<h4><code><a title="audiocraft.utils.samples.manager.SampleManager" href="#audiocraft.utils.samples.manager.SampleManager">SampleManager</a></code></h4>
|
||
<ul class="">
|
||
<li><code><a title="audiocraft.utils.samples.manager.SampleManager.add_sample" href="#audiocraft.utils.samples.manager.SampleManager.add_sample">add_sample</a></code></li>
|
||
<li><code><a title="audiocraft.utils.samples.manager.SampleManager.add_samples" href="#audiocraft.utils.samples.manager.SampleManager.add_samples">add_samples</a></code></li>
|
||
<li><code><a title="audiocraft.utils.samples.manager.SampleManager.get_samples" href="#audiocraft.utils.samples.manager.SampleManager.get_samples">get_samples</a></code></li>
|
||
<li><code><a title="audiocraft.utils.samples.manager.SampleManager.latest_epoch" href="#audiocraft.utils.samples.manager.SampleManager.latest_epoch">latest_epoch</a></code></li>
|
||
</ul>
|
||
</li>
|
||
</ul>
|
||
</li>
|
||
</ul>
|
||
</nav>
|
||
</main>
|
||
<footer id="footer">
|
||
<p>Generated by <a href="https://pdoc3.github.io/pdoc" title="pdoc: Python API documentation generator"><cite>pdoc</cite> 0.11.5</a>.</p>
|
||
</footer>
|
||
</body>
|
||
</html>
|