facebookresearch--audiocraft
690 行
36 KiB
HTML
690 行
36 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.cache API documentation</title>
|
||
<meta name="description" content="">
|
||
<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.cache</code></h1>
|
||
</header>
|
||
<section id="section-intro">
|
||
</section>
|
||
<section>
|
||
</section>
|
||
<section>
|
||
</section>
|
||
<section>
|
||
<h2 class="section-title" id="header-functions">Functions</h2>
|
||
<dl>
|
||
<dt id="audiocraft.utils.cache.get_full_embed"><code class="name flex">
|
||
<span>def <span class="ident">get_full_embed</span></span>(<span>full_embed: torch.Tensor, x: Any, idx: int, device: torch.device | str) ‑> torch.Tensor</span>
|
||
</code></dt>
|
||
<dd>
|
||
<details class="source">
|
||
<summary>
|
||
<span>Expand source code</span>
|
||
</summary>
|
||
<pre><code class="python">def get_full_embed(full_embed: torch.Tensor, x: tp.Any, idx: int, device: tp.Union[str, torch.device]) -> torch.Tensor:
|
||
"""Utility function for the EmbeddingCache, returning the full embedding without any chunking.
|
||
This method can be used in case there is no need in extracting a chunk of the full embedding
|
||
read from the cache.
|
||
|
||
Args:
|
||
full_embed (torch.Tensor): The full embedding.
|
||
x (any): Batch object from which the full embedding is derived.
|
||
idx (torch.Tensor): Index of object to consider in the batch object.
|
||
Returns:
|
||
full_embed (torch.Tensor): The full embedding
|
||
"""
|
||
return full_embed.to(device)</code></pre>
|
||
</details>
|
||
<div class="desc"><p>Utility function for the EmbeddingCache, returning the full embedding without any chunking.
|
||
This method can be used in case there is no need in extracting a chunk of the full embedding
|
||
read from the cache.</p>
|
||
<h2 id="args">Args</h2>
|
||
<dl>
|
||
<dt><strong><code>full_embed</code></strong> : <code>torch.Tensor</code></dt>
|
||
<dd>The full embedding.</dd>
|
||
<dt><strong><code>x</code></strong> : <code>any</code></dt>
|
||
<dd>Batch object from which the full embedding is derived.</dd>
|
||
<dt><strong><code>idx</code></strong> : <code>torch.Tensor</code></dt>
|
||
<dd>Index of object to consider in the batch object.</dd>
|
||
</dl>
|
||
<h2 id="returns">Returns</h2>
|
||
<p>full_embed (torch.Tensor): The full embedding</p></div>
|
||
</dd>
|
||
</dl>
|
||
</section>
|
||
<section>
|
||
<h2 class="section-title" id="header-classes">Classes</h2>
|
||
<dl>
|
||
<dt id="audiocraft.utils.cache.CachedBatchLoader"><code class="flex name class">
|
||
<span>class <span class="ident">CachedBatchLoader</span></span>
|
||
<span>(</span><span>cache_folder: pathlib.Path,<br>batch_size: int,<br>num_workers: int = 10,<br>min_length: int = 1)</span>
|
||
</code></dt>
|
||
<dd>
|
||
<details class="source">
|
||
<summary>
|
||
<span>Expand source code</span>
|
||
</summary>
|
||
<pre><code class="python">class CachedBatchLoader:
|
||
"""Loader for cached mini-batches dumped with `CachedBatchWriter`.
|
||
|
||
Args:
|
||
cache_folder (Path): folder in which the cached minibatches are stored.
|
||
batch_size (int): batch size (per GPU) expected.
|
||
num_workers (int): number of workers to use for loading.
|
||
min_length (int): minimum expected length for each epoch. If some
|
||
mini-batches are missing, and error is raised.
|
||
|
||
This is iterable just like a regular DataLoader.
|
||
"""
|
||
|
||
def __init__(self, cache_folder: Path, batch_size: int,
|
||
num_workers: int = 10, min_length: int = 1):
|
||
self.cache_folder = cache_folder
|
||
self.batch_size = batch_size
|
||
self.num_workers = num_workers
|
||
self.min_length = min_length
|
||
self._current_epoch: tp.Optional[int] = None
|
||
self.sampler = None # for compatibility with the regular DataLoader
|
||
|
||
def __len__(self):
|
||
path = CachedBatchWriter._get_zip_path(self.cache_folder, self._current_epoch or 0, 0).parent
|
||
return len([p for p in path.iterdir() if p.suffix == ".zip"])
|
||
|
||
def start_epoch(self, epoch: int):
|
||
"""Call at the beginning of each epoch.
|
||
"""
|
||
self._current_epoch = epoch
|
||
|
||
def _zip_path(self, index: int):
|
||
assert self._current_epoch is not None
|
||
return CachedBatchWriter._get_zip_path(self.cache_folder, self._current_epoch, index)
|
||
|
||
def _load_one(self, index: int):
|
||
zip_path = self._zip_path(index)
|
||
if not zip_path.exists():
|
||
if index < self.min_length:
|
||
raise RuntimeError(f"Cache should have at least {self.min_length} batches, but {index} doesn't exist")
|
||
|
||
return None
|
||
mode = "rb" if sys.version_info >= (3, 9) else "r"
|
||
try:
|
||
with zipfile.ZipFile(zip_path, 'r') as zf:
|
||
rank = flashy.distrib.rank()
|
||
world_size = flashy.distrib.world_size()
|
||
root = zipfile.Path(zf)
|
||
items = list(root.iterdir())
|
||
total_batch_size = self.batch_size * world_size
|
||
if len(items) < total_batch_size:
|
||
raise RuntimeError(
|
||
f"The cache can handle a max batch size of {len(items)}, "
|
||
f"but {total_batch_size} is needed.")
|
||
start = rank * self.batch_size
|
||
items = items[start: start + self.batch_size]
|
||
assert len(items) == self.batch_size
|
||
entries = []
|
||
entries = [torch.load(item.open(mode), 'cpu') for item in items] # type: ignore
|
||
transposed = zip(*entries)
|
||
out = []
|
||
for part in transposed:
|
||
assert len(part) > 0
|
||
if isinstance(part[0], torch.Tensor):
|
||
out.append(torch.stack(part))
|
||
else:
|
||
assert isinstance(part, torch.Tensor)
|
||
out.append(part)
|
||
return out
|
||
except Exception:
|
||
logger.error("Error when reading zip path %s", zip_path)
|
||
raise
|
||
|
||
def __iter__(self):
|
||
"""This will yields tuples, exactly as provided to the
|
||
`CachedBatchWriter.save` method.
|
||
"""
|
||
pool = ThreadPoolExecutor(self.num_workers)
|
||
next_index = 0
|
||
queue = deque()
|
||
|
||
def _get_next():
|
||
nonlocal next_index
|
||
r = queue.popleft().result()
|
||
if r is None:
|
||
return None
|
||
else:
|
||
queue.append(pool.submit(self._load_one, next_index))
|
||
next_index += 1
|
||
return r
|
||
|
||
with pool:
|
||
# fill the buffer of fetching jobs.
|
||
for _ in range(2 * self.num_workers):
|
||
queue.append(pool.submit(self._load_one, next_index))
|
||
next_index += 1
|
||
while True:
|
||
batch = _get_next()
|
||
if batch is None:
|
||
return
|
||
yield batch</code></pre>
|
||
</details>
|
||
<div class="desc"><p>Loader for cached mini-batches dumped with <code><a title="audiocraft.utils.cache.CachedBatchWriter" href="#audiocraft.utils.cache.CachedBatchWriter">CachedBatchWriter</a></code>.</p>
|
||
<h2 id="args">Args</h2>
|
||
<dl>
|
||
<dt><strong><code>cache_folder</code></strong> : <code>Path</code></dt>
|
||
<dd>folder in which the cached minibatches are stored.</dd>
|
||
<dt><strong><code>batch_size</code></strong> : <code>int</code></dt>
|
||
<dd>batch size (per GPU) expected.</dd>
|
||
<dt><strong><code>num_workers</code></strong> : <code>int</code></dt>
|
||
<dd>number of workers to use for loading.</dd>
|
||
<dt><strong><code>min_length</code></strong> : <code>int</code></dt>
|
||
<dd>minimum expected length for each epoch. If some
|
||
mini-batches are missing, and error is raised.</dd>
|
||
</dl>
|
||
<p>This is iterable just like a regular DataLoader.</p></div>
|
||
<h3>Methods</h3>
|
||
<dl>
|
||
<dt id="audiocraft.utils.cache.CachedBatchLoader.start_epoch"><code class="name flex">
|
||
<span>def <span class="ident">start_epoch</span></span>(<span>self, epoch: int)</span>
|
||
</code></dt>
|
||
<dd>
|
||
<details class="source">
|
||
<summary>
|
||
<span>Expand source code</span>
|
||
</summary>
|
||
<pre><code class="python">def start_epoch(self, epoch: int):
|
||
"""Call at the beginning of each epoch.
|
||
"""
|
||
self._current_epoch = epoch</code></pre>
|
||
</details>
|
||
<div class="desc"><p>Call at the beginning of each epoch.</p></div>
|
||
</dd>
|
||
</dl>
|
||
</dd>
|
||
<dt id="audiocraft.utils.cache.CachedBatchWriter"><code class="flex name class">
|
||
<span>class <span class="ident">CachedBatchWriter</span></span>
|
||
<span>(</span><span>cache_folder: pathlib.Path)</span>
|
||
</code></dt>
|
||
<dd>
|
||
<details class="source">
|
||
<summary>
|
||
<span>Expand source code</span>
|
||
</summary>
|
||
<pre><code class="python">class CachedBatchWriter:
|
||
"""Write pre computed caches for mini batches. This can
|
||
make loading a lot more efficient depending on your filesystem.
|
||
|
||
Args:
|
||
cache_folder (Path): folder in which the cached minibatches
|
||
will be stored.
|
||
|
||
Inside cache folder, the structure is the following:
|
||
`epoch_number / update_number.zip`
|
||
And the zip file contains one entry per batch item.
|
||
|
||
It is possible to use the cache with a batch size smaller than
|
||
created with but obviously not larger. Make sure to call the
|
||
`start_epoch(epoch)` method for indicating changes of epochs.
|
||
|
||
See the grid `audiocraft/grids/musicgen/musicgen_warmup_cache.py`
|
||
for an example of how to warmup the cache.
|
||
"""
|
||
def __init__(self, cache_folder: Path):
|
||
self.cache_folder = cache_folder
|
||
self._current_epoch: tp.Optional[int] = None
|
||
self._current_index = 0
|
||
|
||
def start_epoch(self, epoch: int):
|
||
"""Call at the beginning of each epoch.
|
||
"""
|
||
self._current_epoch = epoch
|
||
self._current_index = 0
|
||
self._zip_path.parent.mkdir(exist_ok=True, parents=True)
|
||
|
||
@staticmethod
|
||
def _get_zip_path(cache_folder: Path, epoch: int, index: int):
|
||
return cache_folder / f"{epoch:05d}" / f"{index:06d}.zip"
|
||
|
||
@property
|
||
def _zip_path(self):
|
||
assert self._current_epoch is not None
|
||
return CachedBatchWriter._get_zip_path(self.cache_folder, self._current_epoch, self._current_index)
|
||
|
||
def save(self, *content):
|
||
"""Save one mini batch. This function is distributed-aware
|
||
and will automatically merge all the items from the different
|
||
workers.
|
||
"""
|
||
all_contents = []
|
||
for rank in range(flashy.distrib.world_size()):
|
||
their_content = flashy.distrib.broadcast_object(content, src=rank)
|
||
all_contents.append(their_content)
|
||
|
||
if flashy.distrib.is_rank_zero():
|
||
idx = 0
|
||
with flashy.utils.write_and_rename(self._zip_path) as tmp:
|
||
with zipfile.ZipFile(tmp, 'w') as zf:
|
||
for content in all_contents:
|
||
for vals in zip(*content):
|
||
with zf.open(f'{idx}', 'w') as f: # type: ignore
|
||
torch.save(vals, f)
|
||
idx += 1
|
||
flashy.distrib.barrier()
|
||
self._current_index += 1</code></pre>
|
||
</details>
|
||
<div class="desc"><p>Write pre computed caches for mini batches. This can
|
||
make loading a lot more efficient depending on your filesystem.</p>
|
||
<h2 id="args">Args</h2>
|
||
<dl>
|
||
<dt><strong><code>cache_folder</code></strong> : <code>Path</code></dt>
|
||
<dd>folder in which the cached minibatches
|
||
will be stored.</dd>
|
||
</dl>
|
||
<p>Inside cache folder, the structure is the following:
|
||
<code>epoch_number / update_number.zip</code>
|
||
And the zip file contains one entry per batch item.</p>
|
||
<p>It is possible to use the cache with a batch size smaller than
|
||
created with but obviously not larger. Make sure to call the
|
||
<code>start_epoch(epoch)</code> method for indicating changes of epochs.</p>
|
||
<p>See the grid <code>audiocraft/grids/musicgen/musicgen_warmup_cache.py</code>
|
||
for an example of how to warmup the cache.</p></div>
|
||
<h3>Methods</h3>
|
||
<dl>
|
||
<dt id="audiocraft.utils.cache.CachedBatchWriter.save"><code class="name flex">
|
||
<span>def <span class="ident">save</span></span>(<span>self, *content)</span>
|
||
</code></dt>
|
||
<dd>
|
||
<details class="source">
|
||
<summary>
|
||
<span>Expand source code</span>
|
||
</summary>
|
||
<pre><code class="python">def save(self, *content):
|
||
"""Save one mini batch. This function is distributed-aware
|
||
and will automatically merge all the items from the different
|
||
workers.
|
||
"""
|
||
all_contents = []
|
||
for rank in range(flashy.distrib.world_size()):
|
||
their_content = flashy.distrib.broadcast_object(content, src=rank)
|
||
all_contents.append(their_content)
|
||
|
||
if flashy.distrib.is_rank_zero():
|
||
idx = 0
|
||
with flashy.utils.write_and_rename(self._zip_path) as tmp:
|
||
with zipfile.ZipFile(tmp, 'w') as zf:
|
||
for content in all_contents:
|
||
for vals in zip(*content):
|
||
with zf.open(f'{idx}', 'w') as f: # type: ignore
|
||
torch.save(vals, f)
|
||
idx += 1
|
||
flashy.distrib.barrier()
|
||
self._current_index += 1</code></pre>
|
||
</details>
|
||
<div class="desc"><p>Save one mini batch. This function is distributed-aware
|
||
and will automatically merge all the items from the different
|
||
workers.</p></div>
|
||
</dd>
|
||
<dt id="audiocraft.utils.cache.CachedBatchWriter.start_epoch"><code class="name flex">
|
||
<span>def <span class="ident">start_epoch</span></span>(<span>self, epoch: int)</span>
|
||
</code></dt>
|
||
<dd>
|
||
<details class="source">
|
||
<summary>
|
||
<span>Expand source code</span>
|
||
</summary>
|
||
<pre><code class="python">def start_epoch(self, epoch: int):
|
||
"""Call at the beginning of each epoch.
|
||
"""
|
||
self._current_epoch = epoch
|
||
self._current_index = 0
|
||
self._zip_path.parent.mkdir(exist_ok=True, parents=True)</code></pre>
|
||
</details>
|
||
<div class="desc"><p>Call at the beginning of each epoch.</p></div>
|
||
</dd>
|
||
</dl>
|
||
</dd>
|
||
<dt id="audiocraft.utils.cache.EmbeddingCache"><code class="flex name class">
|
||
<span>class <span class="ident">EmbeddingCache</span></span>
|
||
<span>(</span><span>cache_path: str | pathlib.Path,<br>device: torch.device | str,<br>compute_embed_fn: Callable[[pathlib.Path, Any, int], torch.Tensor],<br>extract_embed_fn: Callable[[torch.Tensor, Any, int], torch.Tensor] | None = None)</span>
|
||
</code></dt>
|
||
<dd>
|
||
<details class="source">
|
||
<summary>
|
||
<span>Expand source code</span>
|
||
</summary>
|
||
<pre><code class="python">class EmbeddingCache:
|
||
"""Cache around embeddings computation for faster execution.
|
||
The EmbeddingCache is storing pre-computed embeddings on disk and provides a simple API
|
||
to retrieve the pre-computed embeddings on full inputs and extract only a given chunk
|
||
using a user-provided function. When the cache is warm (all embeddings are pre-computed),
|
||
the EmbeddingCache allows for faster training as it removes the need of computing the embeddings.
|
||
Additionally, it provides in-memory cache around the loaded embeddings to limit IO footprint
|
||
and synchronization points in the forward calls.
|
||
|
||
Args:
|
||
cache_path (Path): Path to folder where all pre-computed embeddings are saved on disk.
|
||
device (str or torch.device): Device on which the embedding is returned.
|
||
compute_embed_fn (callable[[Path, any, int], torch.Tensor], optional): Function to compute
|
||
the embedding from a given object and path. This user provided function can compute the
|
||
embedding from the provided object or using the provided path as entry point. The last parameter
|
||
specify the index corresponding to the current embedding in the object that can represent batch metadata.
|
||
extract_embed_fn (callable[[torch.Tensor, any, int], torch.Tensor], optional): Function to extract
|
||
the desired embedding chunk from the full embedding loaded from the cache. The last parameter
|
||
specify the index corresponding to the current embedding in the object that can represent batch metadata.
|
||
If not specified, will return the full embedding unmodified.
|
||
"""
|
||
def __init__(self, cache_path: tp.Union[str, Path], device: tp.Union[str, torch.device],
|
||
compute_embed_fn: tp.Callable[[Path, tp.Any, int], torch.Tensor],
|
||
extract_embed_fn: tp.Optional[tp.Callable[[torch.Tensor, tp.Any, int], torch.Tensor]] = None):
|
||
self.cache_path = Path(cache_path)
|
||
self.device = device
|
||
self._compute_embed_fn = compute_embed_fn
|
||
self._extract_embed_fn: tp.Callable[[torch.Tensor, tp.Any, int], torch.Tensor]
|
||
if extract_embed_fn is not None:
|
||
self._extract_embed_fn = extract_embed_fn
|
||
else:
|
||
self._extract_embed_fn = partial(get_full_embed, device=device)
|
||
if self.cache_path is not None:
|
||
self.cache_path.mkdir(exist_ok=True, parents=True)
|
||
logger.info(f"Cache instantiated at: {self.cache_path}")
|
||
self.pool = ThreadPoolExecutor(8)
|
||
self.pool.__enter__()
|
||
self._current_batch_cache: dict = {}
|
||
self._memory_cache: dict = {}
|
||
|
||
def _get_cache_path(self, path: tp.Union[Path, str]):
|
||
"""Get cache path for the given file path."""
|
||
sig = sha1(str(path).encode()).hexdigest()
|
||
return self.cache_path / sig
|
||
|
||
@staticmethod
|
||
def _get_full_embed_from_cache(cache: Path):
|
||
"""Loads full pre-computed embedding from the cache."""
|
||
try:
|
||
embed = torch.load(cache, 'cpu')
|
||
except Exception as exc:
|
||
logger.error("Error loading %s: %r", cache, exc)
|
||
embed = None
|
||
return embed
|
||
|
||
def get_embed_from_cache(self, paths: tp.List[Path], x: tp.Any) -> torch.Tensor:
|
||
"""Get embedding from cache, computing and storing it to cache if not already cached.
|
||
The EmbeddingCache first tries to load the embedding from the in-memory cache
|
||
containing the pre-computed chunks populated through `populate_embed_cache`.
|
||
If not found, the full embedding is computed and stored on disk to be later accessed
|
||
to populate the in-memory cache, and the desired embedding chunk is extracted and returned.
|
||
|
||
Args:
|
||
paths (list[Path or str]): List of paths from where the embeddings can be loaded.
|
||
x (any): Object from which the embedding is extracted.
|
||
"""
|
||
embeds = []
|
||
for idx, path in enumerate(paths):
|
||
cache = self._get_cache_path(path)
|
||
if cache in self._current_batch_cache:
|
||
embed = self._current_batch_cache[cache]
|
||
else:
|
||
full_embed = self._compute_embed_fn(path, x, idx)
|
||
try:
|
||
with flashy.utils.write_and_rename(cache, pid=True) as f:
|
||
torch.save(full_embed.cpu(), f)
|
||
except Exception as exc:
|
||
logger.error('Error saving embed %s (%s): %r', cache, full_embed.shape, exc)
|
||
else:
|
||
logger.info('New embed cache saved: %s (%s)', cache, full_embed.shape)
|
||
embed = self._extract_embed_fn(full_embed, x, idx)
|
||
embeds.append(embed)
|
||
embed = torch.stack(embeds, dim=0)
|
||
return embed
|
||
|
||
def populate_embed_cache(self, paths: tp.List[Path], x: tp.Any) -> None:
|
||
"""Populate in-memory caches for embeddings reading from the embeddings stored on disk.
|
||
The in-memory caches consist in a cache for the full embedding and another cache for the
|
||
final embedding chunk. Such caches are used to limit the IO access when computing the actual embeddings
|
||
and reduce the IO footprint and synchronization points during forward passes.
|
||
|
||
Args:
|
||
paths (list[Path]): List of paths from where the embeddings can be loaded.
|
||
x (any): Object from which the embedding is extracted.
|
||
"""
|
||
self._current_batch_cache.clear()
|
||
if self.cache_path is not None:
|
||
futures: list = []
|
||
for path in paths:
|
||
assert path is not None, "Path is required for computation from cache"
|
||
cache = self._get_cache_path(path)
|
||
if cache in self._memory_cache or not cache.exists():
|
||
futures.append(None)
|
||
else:
|
||
futures.append(self.pool.submit(EmbeddingCache._get_full_embed_from_cache, cache))
|
||
for idx, (path, future) in enumerate(zip(paths, futures)):
|
||
assert path is not None
|
||
cache = self._get_cache_path(path)
|
||
full_embed = None
|
||
if future is None:
|
||
if cache in self._memory_cache:
|
||
full_embed = self._memory_cache[cache]
|
||
else:
|
||
full_embed = future.result()
|
||
if full_embed is not None:
|
||
self._memory_cache[cache] = full_embed
|
||
full_embed = full_embed.to(self.device)
|
||
if full_embed is not None:
|
||
embed = self._extract_embed_fn(full_embed, x, idx)
|
||
self._current_batch_cache[cache] = embed</code></pre>
|
||
</details>
|
||
<div class="desc"><p>Cache around embeddings computation for faster execution.
|
||
The EmbeddingCache is storing pre-computed embeddings on disk and provides a simple API
|
||
to retrieve the pre-computed embeddings on full inputs and extract only a given chunk
|
||
using a user-provided function. When the cache is warm (all embeddings are pre-computed),
|
||
the EmbeddingCache allows for faster training as it removes the need of computing the embeddings.
|
||
Additionally, it provides in-memory cache around the loaded embeddings to limit IO footprint
|
||
and synchronization points in the forward calls.</p>
|
||
<h2 id="args">Args</h2>
|
||
<dl>
|
||
<dt><strong><code>cache_path</code></strong> : <code>Path</code></dt>
|
||
<dd>Path to folder where all pre-computed embeddings are saved on disk.</dd>
|
||
<dt><strong><code>device</code></strong> : <code>str</code> or <code>torch.device</code></dt>
|
||
<dd>Device on which the embedding is returned.</dd>
|
||
<dt><strong><code>compute_embed_fn</code></strong> : <code>callable[[Path, any, int], torch.Tensor]</code>, optional</dt>
|
||
<dd>Function to compute
|
||
the embedding from a given object and path. This user provided function can compute the
|
||
embedding from the provided object or using the provided path as entry point. The last parameter
|
||
specify the index corresponding to the current embedding in the object that can represent batch metadata.</dd>
|
||
<dt><strong><code>extract_embed_fn</code></strong> : <code>callable[[torch.Tensor, any, int], torch.Tensor]</code>, optional</dt>
|
||
<dd>Function to extract
|
||
the desired embedding chunk from the full embedding loaded from the cache. The last parameter
|
||
specify the index corresponding to the current embedding in the object that can represent batch metadata.
|
||
If not specified, will return the full embedding unmodified.</dd>
|
||
</dl></div>
|
||
<h3>Methods</h3>
|
||
<dl>
|
||
<dt id="audiocraft.utils.cache.EmbeddingCache.get_embed_from_cache"><code class="name flex">
|
||
<span>def <span class="ident">get_embed_from_cache</span></span>(<span>self, paths: List[pathlib.Path], x: Any) ‑> torch.Tensor</span>
|
||
</code></dt>
|
||
<dd>
|
||
<details class="source">
|
||
<summary>
|
||
<span>Expand source code</span>
|
||
</summary>
|
||
<pre><code class="python">def get_embed_from_cache(self, paths: tp.List[Path], x: tp.Any) -> torch.Tensor:
|
||
"""Get embedding from cache, computing and storing it to cache if not already cached.
|
||
The EmbeddingCache first tries to load the embedding from the in-memory cache
|
||
containing the pre-computed chunks populated through `populate_embed_cache`.
|
||
If not found, the full embedding is computed and stored on disk to be later accessed
|
||
to populate the in-memory cache, and the desired embedding chunk is extracted and returned.
|
||
|
||
Args:
|
||
paths (list[Path or str]): List of paths from where the embeddings can be loaded.
|
||
x (any): Object from which the embedding is extracted.
|
||
"""
|
||
embeds = []
|
||
for idx, path in enumerate(paths):
|
||
cache = self._get_cache_path(path)
|
||
if cache in self._current_batch_cache:
|
||
embed = self._current_batch_cache[cache]
|
||
else:
|
||
full_embed = self._compute_embed_fn(path, x, idx)
|
||
try:
|
||
with flashy.utils.write_and_rename(cache, pid=True) as f:
|
||
torch.save(full_embed.cpu(), f)
|
||
except Exception as exc:
|
||
logger.error('Error saving embed %s (%s): %r', cache, full_embed.shape, exc)
|
||
else:
|
||
logger.info('New embed cache saved: %s (%s)', cache, full_embed.shape)
|
||
embed = self._extract_embed_fn(full_embed, x, idx)
|
||
embeds.append(embed)
|
||
embed = torch.stack(embeds, dim=0)
|
||
return embed</code></pre>
|
||
</details>
|
||
<div class="desc"><p>Get embedding from cache, computing and storing it to cache if not already cached.
|
||
The EmbeddingCache first tries to load the embedding from the in-memory cache
|
||
containing the pre-computed chunks populated through <code>populate_embed_cache</code>.
|
||
If not found, the full embedding is computed and stored on disk to be later accessed
|
||
to populate the in-memory cache, and the desired embedding chunk is extracted and returned.</p>
|
||
<h2 id="args">Args</h2>
|
||
<dl>
|
||
<dt><strong><code>paths</code></strong> : <code>list[Path</code> or <code>str]</code></dt>
|
||
<dd>List of paths from where the embeddings can be loaded.</dd>
|
||
<dt><strong><code>x</code></strong> : <code>any</code></dt>
|
||
<dd>Object from which the embedding is extracted.</dd>
|
||
</dl></div>
|
||
</dd>
|
||
<dt id="audiocraft.utils.cache.EmbeddingCache.populate_embed_cache"><code class="name flex">
|
||
<span>def <span class="ident">populate_embed_cache</span></span>(<span>self, paths: List[pathlib.Path], x: Any) ‑> None</span>
|
||
</code></dt>
|
||
<dd>
|
||
<details class="source">
|
||
<summary>
|
||
<span>Expand source code</span>
|
||
</summary>
|
||
<pre><code class="python">def populate_embed_cache(self, paths: tp.List[Path], x: tp.Any) -> None:
|
||
"""Populate in-memory caches for embeddings reading from the embeddings stored on disk.
|
||
The in-memory caches consist in a cache for the full embedding and another cache for the
|
||
final embedding chunk. Such caches are used to limit the IO access when computing the actual embeddings
|
||
and reduce the IO footprint and synchronization points during forward passes.
|
||
|
||
Args:
|
||
paths (list[Path]): List of paths from where the embeddings can be loaded.
|
||
x (any): Object from which the embedding is extracted.
|
||
"""
|
||
self._current_batch_cache.clear()
|
||
if self.cache_path is not None:
|
||
futures: list = []
|
||
for path in paths:
|
||
assert path is not None, "Path is required for computation from cache"
|
||
cache = self._get_cache_path(path)
|
||
if cache in self._memory_cache or not cache.exists():
|
||
futures.append(None)
|
||
else:
|
||
futures.append(self.pool.submit(EmbeddingCache._get_full_embed_from_cache, cache))
|
||
for idx, (path, future) in enumerate(zip(paths, futures)):
|
||
assert path is not None
|
||
cache = self._get_cache_path(path)
|
||
full_embed = None
|
||
if future is None:
|
||
if cache in self._memory_cache:
|
||
full_embed = self._memory_cache[cache]
|
||
else:
|
||
full_embed = future.result()
|
||
if full_embed is not None:
|
||
self._memory_cache[cache] = full_embed
|
||
full_embed = full_embed.to(self.device)
|
||
if full_embed is not None:
|
||
embed = self._extract_embed_fn(full_embed, x, idx)
|
||
self._current_batch_cache[cache] = embed</code></pre>
|
||
</details>
|
||
<div class="desc"><p>Populate in-memory caches for embeddings reading from the embeddings stored on disk.
|
||
The in-memory caches consist in a cache for the full embedding and another cache for the
|
||
final embedding chunk. Such caches are used to limit the IO access when computing the actual embeddings
|
||
and reduce the IO footprint and synchronization points during forward passes.</p>
|
||
<h2 id="args">Args</h2>
|
||
<dl>
|
||
<dt><strong><code>paths</code></strong> : <code>list[Path]</code></dt>
|
||
<dd>List of paths from where the embeddings can be loaded.</dd>
|
||
<dt><strong><code>x</code></strong> : <code>any</code></dt>
|
||
<dd>Object from which the embedding is extracted.</dd>
|
||
</dl></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" href="index.html">audiocraft.utils</a></code></li>
|
||
</ul>
|
||
</li>
|
||
<li><h3><a href="#header-functions">Functions</a></h3>
|
||
<ul class="">
|
||
<li><code><a title="audiocraft.utils.cache.get_full_embed" href="#audiocraft.utils.cache.get_full_embed">get_full_embed</a></code></li>
|
||
</ul>
|
||
</li>
|
||
<li><h3><a href="#header-classes">Classes</a></h3>
|
||
<ul>
|
||
<li>
|
||
<h4><code><a title="audiocraft.utils.cache.CachedBatchLoader" href="#audiocraft.utils.cache.CachedBatchLoader">CachedBatchLoader</a></code></h4>
|
||
<ul class="">
|
||
<li><code><a title="audiocraft.utils.cache.CachedBatchLoader.start_epoch" href="#audiocraft.utils.cache.CachedBatchLoader.start_epoch">start_epoch</a></code></li>
|
||
</ul>
|
||
</li>
|
||
<li>
|
||
<h4><code><a title="audiocraft.utils.cache.CachedBatchWriter" href="#audiocraft.utils.cache.CachedBatchWriter">CachedBatchWriter</a></code></h4>
|
||
<ul class="">
|
||
<li><code><a title="audiocraft.utils.cache.CachedBatchWriter.save" href="#audiocraft.utils.cache.CachedBatchWriter.save">save</a></code></li>
|
||
<li><code><a title="audiocraft.utils.cache.CachedBatchWriter.start_epoch" href="#audiocraft.utils.cache.CachedBatchWriter.start_epoch">start_epoch</a></code></li>
|
||
</ul>
|
||
</li>
|
||
<li>
|
||
<h4><code><a title="audiocraft.utils.cache.EmbeddingCache" href="#audiocraft.utils.cache.EmbeddingCache">EmbeddingCache</a></code></h4>
|
||
<ul class="">
|
||
<li><code><a title="audiocraft.utils.cache.EmbeddingCache.get_embed_from_cache" href="#audiocraft.utils.cache.EmbeddingCache.get_embed_from_cache">get_embed_from_cache</a></code></li>
|
||
<li><code><a title="audiocraft.utils.cache.EmbeddingCache.populate_embed_cache" href="#audiocraft.utils.cache.EmbeddingCache.populate_embed_cache">populate_embed_cache</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>
|