facebookresearch--audiocraft
1680 行
88 KiB
HTML
1680 行
88 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.solvers.base 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.solvers.base</code></h1>
|
||
</header>
|
||
<section id="section-intro">
|
||
</section>
|
||
<section>
|
||
</section>
|
||
<section>
|
||
</section>
|
||
<section>
|
||
</section>
|
||
<section>
|
||
<h2 class="section-title" id="header-classes">Classes</h2>
|
||
<dl>
|
||
<dt id="audiocraft.solvers.base.StandardSolver"><code class="flex name class">
|
||
<span>class <span class="ident">StandardSolver</span></span>
|
||
<span>(</span><span>cfg: omegaconf.dictconfig.DictConfig)</span>
|
||
</code></dt>
|
||
<dd>
|
||
<details class="source">
|
||
<summary>
|
||
<span>Expand source code</span>
|
||
</summary>
|
||
<pre><code class="python">class StandardSolver(ABC, flashy.BaseSolver):
|
||
"""Standard solver for AudioCraft.
|
||
|
||
The standard solver implements a base training loop with the following stages:
|
||
train, valid, evaluate and generate that are expected to be all defined for
|
||
solvers in AudioCraft. It also provides a nice default management of Dora history replay,
|
||
checkpoint management across epoch, and logging configuration.
|
||
|
||
AudioCraft solvers must inherit from the StandardSolver and define the methods
|
||
associated to each stage as well as the show, build_model and build_dataloaders methods.
|
||
"""
|
||
def __init__(self, cfg: omegaconf.DictConfig):
|
||
super().__init__()
|
||
self.logger.info(f"Instantiating solver {self.__class__.__name__} for XP {self.xp.sig}")
|
||
self.logger.info(f"All XP logs are stored in {self.xp.folder}")
|
||
self.cfg = cfg
|
||
self.device = cfg.device
|
||
self.model: nn.Module
|
||
self._continue_best_source_keys = ['best_state', 'fsdp_best_state']
|
||
self._fsdp_modules: tp.List[fsdp.FSDP] = []
|
||
self._ema_sources: nn.ModuleDict = nn.ModuleDict()
|
||
self.ema: tp.Optional[optim.ModuleDictEMA] = None
|
||
self.dataloaders: tp.Dict[str, torch.utils.data.DataLoader] = dict()
|
||
self._log_updates = self.cfg.logging.get('log_updates', 10)
|
||
if self.cfg.logging.log_tensorboard:
|
||
self.init_tensorboard(**self.cfg.get('tensorboard'))
|
||
if self.cfg.logging.log_wandb and self:
|
||
self.init_wandb(**self.cfg.get('wandb'))
|
||
# keep a copy of the best performing state for stateful objects
|
||
# used for evaluation and generation stages
|
||
dtype_best: tp.Optional[torch.dtype] = None
|
||
if self.cfg.fsdp.use:
|
||
dtype_best = getattr(torch, self.cfg.fsdp.param_dtype) # type: ignore
|
||
assert isinstance(dtype_best, torch.dtype)
|
||
elif self.cfg.autocast:
|
||
dtype_best = getattr(torch, self.cfg.autocast_dtype) # type: ignore
|
||
assert isinstance(dtype_best, torch.dtype)
|
||
self.best_state: BestStateDictManager = BestStateDictManager(dtype=dtype_best)
|
||
# Hacky support for keeping a copy of the full best state in rank0.
|
||
self.fsdp_best_state: tp.Dict[str, tp.Any] = {}
|
||
self.register_stateful('best_state', 'fsdp_best_state') # register best_state object to keep it in state_dict
|
||
self._new_best_state: bool = False # should save a new checkpoint
|
||
# instantiate datasets and appropriate number of updates per epoch
|
||
self.build_dataloaders()
|
||
if self.cfg.execute_only is None:
|
||
assert 'train' in self.dataloaders, "The train dataset split must be provided."
|
||
assert 'valid' in self.dataloaders, "The valid dataset split must be provided."
|
||
self.train_updates_per_epoch = len(self.dataloaders['train']) if 'train' in self.dataloaders else 0
|
||
if self.cfg.optim.updates_per_epoch:
|
||
self.train_updates_per_epoch = self.cfg.optim.updates_per_epoch
|
||
self.total_updates = self.train_updates_per_epoch * self.cfg.optim.epochs
|
||
# instantiate model & exponential moving average on the model
|
||
self.build_model()
|
||
self.logger.info("Model hash: %s", model_hash(self.model))
|
||
assert 'model' in self.stateful.sources, \
|
||
"Please register the model to stateful with self.register_stateful('model') in build_model."
|
||
self.profiler = Profiler(self.model, **self.cfg.profiler)
|
||
self.initialize_ema()
|
||
self.register_stateful('ema')
|
||
assert self.ema is None or 'ema' in self.stateful.sources, \
|
||
"Please register the ema to stateful with self.register_stateful('ema') in build_model."
|
||
self.deadlock_detect = DeadlockDetect(**self.cfg.deadlock)
|
||
# basic statistics on the trained model
|
||
model_size = sum(p.numel() for p in self.model.parameters() if p.requires_grad) / 1e6
|
||
# one copy of grad, one copy of momentum, one copy of denominator and model weights.
|
||
# and 4 bytes for each float!
|
||
mem_usage = model_size * 4 * 4 / 1000
|
||
self.logger.info("Model size: %.2f M params", model_size)
|
||
self.logger.info("Base memory usage, with model, grad and optim: %.2f GB", mem_usage)
|
||
|
||
@property
|
||
def autocast(self):
|
||
"""Convenient autocast (or not) using the solver configuration."""
|
||
return TorchAutocast(enabled=self.cfg.autocast, device_type=self.device, dtype=self.autocast_dtype)
|
||
|
||
def _get_state_source(self, name) -> flashy.state.StateDictSource:
|
||
# Internal utility to get a state source from the solver
|
||
return self.stateful.sources[name]
|
||
|
||
@property
|
||
def best_metric_name(self) -> tp.Optional[str]:
|
||
"""Metric name used to identify the best state. This metric should be stored in the metrics
|
||
used on the stage for best state identification (most likely, `valid`). If None, then
|
||
no best state is saved.
|
||
"""
|
||
return None
|
||
|
||
def register_best_state(self, *args: str):
|
||
"""Register state sources in `BestStateDictManager` to keep their best states along with their
|
||
latest states. The best state will be used at evaluation stages instead of the latest states.
|
||
|
||
Shortcut around `BestStateDictManager.register` method. You can pass any number of
|
||
attribute, included nested attributes and those will be included into the checkpoints
|
||
and automatically restored when `BaseSolver.restore` is called.
|
||
"""
|
||
for name in args:
|
||
state_source = self._get_state_source(name)
|
||
assert name in self.stateful.sources, "Registered states in best should be registered in stateful first!"
|
||
self.best_state.register(name, state_source)
|
||
|
||
def register_ema(self, *args: str):
|
||
"""Register state sources for exponential moving average.
|
||
|
||
The registered sources are used to instantiate a ModuleDictEMA instance.
|
||
The ModuleDictEMA keeps a `nn.ModuleDict` module that is updated when self.ema.step() is called
|
||
and swapped with the original state sources with self.swap_ema_state() method.
|
||
|
||
Usage:
|
||
self.register_ema('model')
|
||
"""
|
||
assert self.ema is None, "Cannot register state source to already instantiated EMA."
|
||
for name in args:
|
||
self._ema_sources[name] = getattr(self, name)
|
||
|
||
def wrap_with_fsdp(self, model: torch.nn.Module, *args, **kwargs):
|
||
model = fsdp.wrap_with_fsdp(self.cfg.fsdp, model, *args, **kwargs)
|
||
if isinstance(model, fsdp.FSDP):
|
||
self._fsdp_modules.append(model)
|
||
return model
|
||
|
||
def update_best_state_from_stage(self, stage_name: str = 'valid'):
|
||
"""Update latest best state based on pending metrics of a given stage. This method relies
|
||
on the `BestStateDictManager.update` method to update the best state_dict with latest weights
|
||
if the registered states happen to match to the best performing setup.
|
||
"""
|
||
if self.best_metric_name is None:
|
||
# when no best metric is defined, the last state is always the best
|
||
self._new_best_state = True
|
||
self.logger.info("Updating best state with current state.")
|
||
else:
|
||
assert stage_name in self._pending_metrics, f"Metrics for stage {stage_name} not found."
|
||
assert self.best_metric_name in self._pending_metrics[stage_name], \
|
||
f"Best metric not found in {stage_name} metrics. Cannot register best state"
|
||
current_score = self._pending_metrics[stage_name][self.best_metric_name]
|
||
all_best_metric_scores = [
|
||
past_metrics[stage_name][self.best_metric_name]
|
||
for past_metrics in self.history
|
||
]
|
||
all_best_metric_scores.append(current_score)
|
||
best_score = min(all_best_metric_scores)
|
||
self._new_best_state = current_score == best_score
|
||
if self._new_best_state:
|
||
old_best = min(all_best_metric_scores[:-1] + [float('inf')])
|
||
self.logger.info(
|
||
f"New best state with {self.best_metric_name}={current_score:.3f} (was {old_best:.3f})")
|
||
|
||
if self._new_best_state:
|
||
if self.cfg.fsdp.use:
|
||
# this will give an empty state dict on all ranks but the rank 0
|
||
# which will have a copy in memory of the full model.
|
||
with fsdp.switch_to_full_state_dict(self._fsdp_modules):
|
||
for name in self.best_state.states.keys():
|
||
state_source = self._get_state_source(name)
|
||
self.best_state.update(name, state_source)
|
||
# we save to a different dict.
|
||
self.fsdp_best_state.update(self.best_state.state_dict())
|
||
# We cannot efficiently load fsdp_best_state when using FSDP,
|
||
# so we have do do a second pass, with the local shards.
|
||
for name in self.best_state.states.keys():
|
||
state_source = self._get_state_source(name)
|
||
self.best_state.update(name, state_source)
|
||
|
||
def _load_new_state_dict(self, state_dict: dict) -> dict:
|
||
old_states = {}
|
||
for name, new_state in state_dict.items():
|
||
state_source = self._get_state_source(name)
|
||
old_states[name] = copy_state(state_source.state_dict())
|
||
state_source.load_state_dict(new_state)
|
||
return old_states
|
||
|
||
@contextmanager
|
||
def swap_best_state(self):
|
||
self.logger.debug(f"Swapping to best state for: {', '.join(self.best_state.state_dict().keys())}")
|
||
old_states = self._load_new_state_dict(self.best_state.state_dict())
|
||
try:
|
||
yield
|
||
finally:
|
||
self.logger.debug("Swapping back from best to original state")
|
||
for name, old_state in old_states.items():
|
||
state_source = self._get_state_source(name)
|
||
state_source.load_state_dict(old_state)
|
||
|
||
@contextmanager
|
||
def swap_ema_state(self):
|
||
if self.ema is None:
|
||
yield
|
||
else:
|
||
ema_state_dict = self.ema.state_dict()['state']
|
||
self.logger.debug(f"Swapping to EMA state for: {', '.join(ema_state_dict.keys())}")
|
||
old_states = self._load_new_state_dict(ema_state_dict)
|
||
try:
|
||
yield
|
||
finally:
|
||
self.logger.debug("Swapping back from EMA state to original state")
|
||
for name, old_state in old_states.items():
|
||
state_source = self._get_state_source(name)
|
||
state_source.load_state_dict(old_state)
|
||
|
||
@property
|
||
def is_training(self):
|
||
return self.current_stage == 'train'
|
||
|
||
def log_model_summary(self, model: nn.Module):
|
||
"""Log model summary, architecture and size of the model."""
|
||
self.logger.info(model)
|
||
mb = sum(p.numel() for p in model.parameters()) * 4 / 2 ** 20
|
||
self.logger.info("Size: %.1f MB", mb)
|
||
|
||
@abstractmethod
|
||
def build_model(self):
|
||
"""Method to implement to initialize model."""
|
||
...
|
||
|
||
def initialize_ema(self):
|
||
"""Initialize exponential moving average with the registered sources.
|
||
EMA object is created if the optim.ema.model.decay value is non-null.
|
||
"""
|
||
from .builders import get_ema
|
||
self.ema = get_ema(self._ema_sources, self.cfg.optim.ema)
|
||
if self.ema is None:
|
||
self.logger.info('No EMA on the model.')
|
||
else:
|
||
assert self.cfg.optim.ema.updates > 0
|
||
self.logger.info(
|
||
f'Initializing EMA on the model with decay = {self.ema.decay}'
|
||
f' every {self.cfg.optim.ema.updates} updates'
|
||
)
|
||
|
||
@abstractmethod
|
||
def build_dataloaders(self):
|
||
"""Method to implement to initialize dataloaders."""
|
||
...
|
||
|
||
@abstractmethod
|
||
def show(self):
|
||
"""Method to log any information without running the job."""
|
||
...
|
||
|
||
@property
|
||
def log_updates(self):
|
||
# convenient access to log updates
|
||
return self._log_updates
|
||
|
||
def checkpoint_path(self, **kwargs):
|
||
kwargs.setdefault('use_fsdp', self.cfg.fsdp.use)
|
||
return self.folder / checkpoint.checkpoint_name(**kwargs)
|
||
|
||
def epoch_checkpoint_path(self, epoch: int, **kwargs):
|
||
kwargs.setdefault('use_fsdp', self.cfg.fsdp.use)
|
||
return self.folder / checkpoint.checkpoint_name(str(epoch), **kwargs)
|
||
|
||
def checkpoint_path_with_name(self, name: str, **kwargs):
|
||
kwargs.setdefault('use_fsdp', self.cfg.fsdp.use)
|
||
return self.folder / checkpoint.checkpoint_name(name=name, **kwargs)
|
||
|
||
def save_checkpoints(self):
|
||
"""Save checkpoint, optionally keeping a copy for a given epoch."""
|
||
is_sharded = self.cfg.fsdp.use
|
||
if not flashy.distrib.is_rank_zero() and not is_sharded:
|
||
return
|
||
self.logger.info("Model hash: %s", model_hash(self.model))
|
||
state = self.state_dict()
|
||
epoch = self.epoch - 1 # pushing metrics will increase the epoch in Flashy, so we do -1 here
|
||
|
||
# save minimal state_dict as new checkpoint every X epoch
|
||
if self.cfg.checkpoint.save_every:
|
||
if epoch % self.cfg.checkpoint.save_every == 0:
|
||
minimal_state = state
|
||
if self.cfg.checkpoint.keep_every_states is not None and len(self.cfg.checkpoint.keep_every_states) > 0:
|
||
minimal_state = {
|
||
name: source for name, source in state.items()
|
||
if name in self.cfg.checkpoint.keep_every_states
|
||
}
|
||
epoch_checkpoint_path = self.epoch_checkpoint_path(epoch)
|
||
checkpoint.save_checkpoint(minimal_state, epoch_checkpoint_path, is_sharded)
|
||
|
||
# save checkpoint as latest checkpoint
|
||
if self.cfg.checkpoint.save_last:
|
||
last_checkpoint_path = self.checkpoint_path()
|
||
checkpoint.save_checkpoint(state, last_checkpoint_path, is_sharded)
|
||
|
||
# flush any stale checkpoint to reduce disk footprint
|
||
checkpoint.flush_stale_checkpoints(self.checkpoint_path())
|
||
|
||
def load_from_pretrained(self, name: str) -> dict:
|
||
raise NotImplementedError("Solver does not provide a way to load pretrained models.")
|
||
|
||
def load_checkpoints(self, load_best: bool = False, ignore_state_keys: tp.List[str] = []) -> tp.Optional[dict]:
|
||
"""Load last checkpoint or the one specified in continue_from.
|
||
|
||
Args:
|
||
load_best (bool): Whether to load from best state dict or not.
|
||
Best state dict is always used when not loading the current xp.
|
||
ignore_state_keys (list of str): List of sources to ignore when loading the state, e.g. `optimizer`.
|
||
Returns:
|
||
state (dict, optional): The loaded state dictionary.
|
||
"""
|
||
# load checkpoints from xp folder or cfg.continue_from
|
||
is_sharded = self.cfg.fsdp.use
|
||
load_from_path: tp.Optional[Path] = None
|
||
checkpoint_source: tp.Optional[checkpoint.CheckpointSource] = None
|
||
|
||
if load_best:
|
||
self.logger.info("Trying to load state_dict from best state.")
|
||
|
||
state: tp.Optional[dict] = None
|
||
rank0_checkpoint_path = self.checkpoint_path(use_fsdp=False)
|
||
current_checkpoint_path = self.checkpoint_path()
|
||
_pretrained_prefix = '//pretrained/'
|
||
continue_pretrained = (self.cfg.continue_from or '').startswith(_pretrained_prefix)
|
||
if rank0_checkpoint_path.exists():
|
||
self.logger.info(f"Loading existing checkpoint: {current_checkpoint_path}")
|
||
load_from_path = current_checkpoint_path
|
||
checkpoint.check_sharded_checkpoint(current_checkpoint_path, rank0_checkpoint_path)
|
||
checkpoint_source = checkpoint.CheckpointSource.CURRENT_XP
|
||
elif self.cfg.continue_from and not continue_pretrained:
|
||
self.logger.info(f"Continuing from provided checkpoint: {self.cfg.continue_from}")
|
||
# we're always continuing from consolidated checkpoints: self.cfg.use_fsdp and not continue_best
|
||
load_from_path = checkpoint.resolve_checkpoint_path(self.cfg.continue_from, use_fsdp=False)
|
||
if load_from_path is None:
|
||
self.logger.error('Could not resolve the continue_from checkpoint %s', self.cfg.continue_from)
|
||
raise RuntimeError(f'Could not resolve continue_from checkpoint {self.cfg.continue_from}')
|
||
checkpoint_source = checkpoint.CheckpointSource.OTHER
|
||
|
||
if load_from_path is not None:
|
||
state = checkpoint.load_checkpoint(load_from_path, is_sharded)
|
||
elif continue_pretrained:
|
||
self.logger.info("Loading a pretrained model. Ignoring 'load_best' and 'ignore_state_keys' params.")
|
||
state = self.load_from_pretrained(self.cfg.continue_from[len(_pretrained_prefix):])
|
||
checkpoint_source = checkpoint.CheckpointSource.PRETRAINED
|
||
load_best = True
|
||
|
||
# checkpoints are not from the current xp, we only retrieve the best state
|
||
if checkpoint_source is not None and checkpoint_source != checkpoint.CheckpointSource.CURRENT_XP:
|
||
assert state is not None
|
||
self.logger.info("Checkpoint source is not the current xp: Load state_dict from best state.")
|
||
load_best = True
|
||
state = {key: state[key] for key in self._continue_best_source_keys if key in state}
|
||
# loaded checkpoints are FSDP checkpoints: we're reading the best state
|
||
# from FSDP and we drop the regular best_state
|
||
if 'fsdp_best_state' in state and state['fsdp_best_state']:
|
||
state.pop('best_state', None)
|
||
self.logger.info("... Loaded checkpoint has FSDP best state")
|
||
# FSDP is enabled in the solver, if the loaded checkpoints do not have FSDP support
|
||
# then we're initializing FSDP best state with the regular best state
|
||
elif self.cfg.fsdp.use:
|
||
if 'fsdp_best_state' not in state or not state['fsdp_best_state']:
|
||
# we swap non-FSDP checkpoints best_state to FSDP-compatible best state
|
||
state['fsdp_best_state'] = state.pop('best_state')
|
||
self.logger.info("... Loaded checkpoint does not have FSDP best state. Use regular best state")
|
||
|
||
if state is not None:
|
||
if load_best:
|
||
self.logger.info("Ignoring keys when loading best %r", ignore_state_keys)
|
||
for key in set(ignore_state_keys):
|
||
if key in state:
|
||
state.pop(key)
|
||
has_best_state = 'best_state' in state or 'fsdp_best_state' in state
|
||
assert has_best_state, ("Trying to load best state but neither 'best_state'",
|
||
" or 'fsdp_best_state' found in checkpoints.")
|
||
self.load_state_dict(state)
|
||
|
||
# for FSDP, let's make extra sure nothing bad happened with out of sync
|
||
# checkpoints across workers.
|
||
epoch = float(self.epoch)
|
||
avg_epoch = flashy.distrib.average_metrics({'epoch': epoch})['epoch']
|
||
if avg_epoch != epoch:
|
||
raise RuntimeError(
|
||
f"Inconsistent loading of checkpoints happened, our epoch is {epoch} "
|
||
f"but average of epochs is {avg_epoch}, at least one gpu must have a "
|
||
"different epoch number.")
|
||
|
||
# on load_best, properly reinitialize state_dict, best states and ema
|
||
# otherwise we load from the current xp and don't alter anything
|
||
if load_best:
|
||
self.logger.info("Loading state_dict from best state.")
|
||
if not self.cfg.fsdp.use and self.fsdp_best_state:
|
||
# loading from an FSDP checkpoint but with FSDP deactivated
|
||
self.logger.info("... Loading from FSDP best state dict.")
|
||
self.best_state.load_state_dict(self.fsdp_best_state)
|
||
|
||
# if load_best, we permanently override the regular state_dict with the best state
|
||
if self.cfg.fsdp.use:
|
||
self.logger.info("FSDP is used, loading from FSDP best state.")
|
||
with fsdp.switch_to_full_state_dict(self._fsdp_modules):
|
||
# this might be really fragile but okay for now.
|
||
self.load_state_dict(self.fsdp_best_state)
|
||
else:
|
||
# we permanently swap the stateful objects to their best state
|
||
self._load_new_state_dict(self.best_state.state_dict())
|
||
|
||
# the EMA modules should also be instantiated with best state.
|
||
# the easiest way to do so is to reinitialize a new EMA with best state loaded.
|
||
if self.ema is not None:
|
||
self.logger.info("Re-initializing EMA from best state")
|
||
self.initialize_ema()
|
||
|
||
if self.cfg.fsdp.use:
|
||
self.logger.info("Re-initializing best state after using FSDP best state.")
|
||
for name in self.best_state.states.keys():
|
||
state_source = self._get_state_source(name)
|
||
self.best_state.update(name, state_source)
|
||
|
||
return state
|
||
|
||
def restore(self, load_best: bool = False, replay_metrics: bool = False,
|
||
ignore_state_keys: tp.List[str] = []) -> bool:
|
||
"""Restore the status of a solver for a given xp.
|
||
|
||
Args:
|
||
load_best (bool): if `True`, load the best state from the checkpoint.
|
||
replay_metrics (bool): if `True`, logs all the metrics from past epochs.
|
||
ignore_state_keys (list of str): list of sources to ignore when loading the state, e.g. `optimizer`.
|
||
"""
|
||
self.logger.info("Restoring weights and history.")
|
||
restored_checkpoints = self.load_checkpoints(load_best, ignore_state_keys)
|
||
|
||
self.logger.info("Model hash: %s", model_hash(self.model))
|
||
|
||
if replay_metrics and len(self.history) > 0:
|
||
self.logger.info("Replaying past metrics...")
|
||
for epoch, stages in enumerate(self.history):
|
||
for stage_name, metrics in stages.items():
|
||
# We manually log the metrics summary to the result logger
|
||
# as we don't want to add them to the pending metrics
|
||
self.result_logger._log_summary(stage_name, metrics, step=epoch + 1, step_name='epoch',
|
||
formatter=self.get_formatter(stage_name))
|
||
return restored_checkpoints is not None
|
||
|
||
def commit(self, save_checkpoints: bool = True):
|
||
"""Commit metrics to dora and save checkpoints at the end of an epoch."""
|
||
# we override commit to introduce more complex checkpoint saving behaviors
|
||
self.history.append(self._pending_metrics) # This will increase self.epoch
|
||
if save_checkpoints:
|
||
self.save_checkpoints()
|
||
self._start_epoch()
|
||
if flashy.distrib.is_rank_zero():
|
||
self.xp.link.update_history(self.history)
|
||
|
||
def run_epoch(self):
|
||
"""Run a single epoch with all stages.
|
||
|
||
Metrics for a given stage are stored in _pending_metrics and committed by the solver afterwards.
|
||
Children solvers can extend this method with custom behavior, e.g.:
|
||
|
||
def run_epoch(self):
|
||
... # custom code
|
||
super().run_epoch()
|
||
... # custom code
|
||
"""
|
||
self.run_stage('train', self.train)
|
||
with torch.no_grad():
|
||
with self.swap_ema_state():
|
||
self.run_stage('valid', self.valid)
|
||
# the best state is updated with EMA states if available
|
||
self.update_best_state_from_stage('valid')
|
||
with self.swap_best_state():
|
||
if self.should_run_stage('evaluate'):
|
||
self.run_stage('evaluate', self.evaluate)
|
||
if self.should_run_stage('generate'):
|
||
self.run_stage('generate', with_rank_rng()(self.generate))
|
||
|
||
def run(self):
|
||
"""Training loop."""
|
||
assert len(self.state_dict()) > 0
|
||
self.restore(replay_metrics=True) # load checkpoint and replay history
|
||
self.log_hyperparams(dict_from_config(self.cfg))
|
||
for epoch in range(self.epoch, self.cfg.optim.epochs + 1):
|
||
if self.should_stop_training():
|
||
return
|
||
self.run_epoch()
|
||
# Commit will send the metrics to Dora and save checkpoints by default.
|
||
self.commit()
|
||
|
||
def should_stop_training(self) -> bool:
|
||
"""Check whether we should stop training or not."""
|
||
return self.epoch > self.cfg.optim.epochs
|
||
|
||
def should_run_stage(self, stage_name) -> bool:
|
||
"""Check whether we want to run the specified stages."""
|
||
stage_every = self.cfg[stage_name].get('every', None)
|
||
is_last_epoch = self.epoch == self.cfg.optim.epochs
|
||
is_epoch_every = (stage_every and self.epoch % stage_every == 0)
|
||
return is_last_epoch or is_epoch_every
|
||
|
||
@abstractmethod
|
||
def run_step(self, idx: int, batch: tp.Any, metrics: dict):
|
||
"""Perform one training or valid step on a given batch."""
|
||
...
|
||
|
||
def common_train_valid(self, dataset_split: str, **kwargs: tp.Any):
|
||
"""Common logic for train and valid stages."""
|
||
self.model.train(self.is_training)
|
||
|
||
loader = self.dataloaders[dataset_split]
|
||
# get a different order for distributed training, otherwise this will get ignored
|
||
if flashy.distrib.world_size() > 1 \
|
||
and isinstance(loader.sampler, torch.utils.data.distributed.DistributedSampler):
|
||
loader.sampler.set_epoch(self.epoch)
|
||
updates_per_epoch = self.train_updates_per_epoch if self.is_training else len(loader)
|
||
if self.cfg.benchmark_no_load:
|
||
self.logger.warning("Fake loading for benchmarking: re-using first batch")
|
||
batch = next(iter(loader))
|
||
loader = [batch] * updates_per_epoch # type: ignore
|
||
lp = self.log_progress(self.current_stage, loader, total=updates_per_epoch, updates=self.log_updates)
|
||
average = flashy.averager() # epoch wise average
|
||
instant_average = flashy.averager() # average between two logging
|
||
metrics: dict = {}
|
||
|
||
with self.profiler, self.deadlock_detect: # profiler will only run for the first 20 updates.
|
||
for idx, batch in enumerate(lp):
|
||
self.deadlock_detect.update('batch')
|
||
if idx >= updates_per_epoch:
|
||
break
|
||
metrics = {}
|
||
metrics = self.run_step(idx, batch, metrics)
|
||
self.deadlock_detect.update('step')
|
||
# run EMA step
|
||
if self.ema is not None and self.is_training and (idx + 1) % self.cfg.optim.ema.updates == 0:
|
||
self.logger.debug("EMA model step")
|
||
self.ema.step()
|
||
self.deadlock_detect.update('ema')
|
||
self.profiler.step()
|
||
instant_metrics = instant_average(metrics)
|
||
if lp.update(**instant_metrics):
|
||
instant_average = flashy.averager() # reset averager between two logging
|
||
metrics = average(metrics) # epoch wise average
|
||
self.deadlock_detect.update('end_batch')
|
||
|
||
metrics = flashy.distrib.average_metrics(metrics, updates_per_epoch)
|
||
return metrics
|
||
|
||
def train(self):
|
||
"""Train stage."""
|
||
return self.common_train_valid('train')
|
||
|
||
def valid(self):
|
||
"""Valid stage."""
|
||
return self.common_train_valid('valid')
|
||
|
||
@abstractmethod
|
||
def evaluate(self):
|
||
"""Evaluate stage."""
|
||
...
|
||
|
||
@abstractmethod
|
||
def generate(self):
|
||
"""Generate stage."""
|
||
...
|
||
|
||
def run_one_stage(self, stage_name: str):
|
||
"""Run only the specified stage.
|
||
This method is useful to only generate samples from a trained experiment
|
||
or rerun the validation or evaluation stages.
|
||
"""
|
||
fn = {
|
||
'generate': with_rank_rng()(self.generate),
|
||
'evaluate': self.evaluate,
|
||
'valid': self.valid,
|
||
}
|
||
if stage_name not in fn:
|
||
raise ValueError(f'Trying to run stage {stage_name} is not supported.')
|
||
assert len(self.state_dict()) > 0
|
||
self._start_epoch()
|
||
with torch.no_grad(), self.swap_best_state():
|
||
self.run_stage(stage_name, fn[stage_name])
|
||
if not self.cfg.execute_inplace:
|
||
self.commit(save_checkpoints=False)
|
||
|
||
@staticmethod
|
||
def get_eval_solver_from_sig(sig: str, dtype: tp.Optional[str] = None,
|
||
device: tp.Optional[str] = None, autocast: bool = True,
|
||
batch_size: tp.Optional[int] = None,
|
||
override_cfg: tp.Optional[tp.Union[dict, omegaconf.DictConfig]] = None,
|
||
**kwargs):
|
||
"""Mostly a convenience function around audiocraft.train.get_solver_from_sig,
|
||
populating all the proper param, deactivating EMA, FSDP, loading the best state,
|
||
basically all you need to get a solver ready to "play" with in single GPU mode
|
||
and with minimal memory overhead.
|
||
|
||
Args:
|
||
sig (str): signature to load.
|
||
dtype (str or None): potential dtype, as a string, i.e. 'float16'.
|
||
device (str or None): potential device, as a string, i.e. 'cuda'.
|
||
override_cfg (dict or omegaconf.DictConfig or None): potential device, as a string, i.e. 'cuda'.
|
||
"""
|
||
from audiocraft import train
|
||
our_override_cfg: tp.Dict[str, tp.Any] = {'optim': {'ema': {'use': False}}}
|
||
our_override_cfg['autocast'] = autocast
|
||
if dtype is not None:
|
||
our_override_cfg['dtype'] = dtype
|
||
if device is not None:
|
||
our_override_cfg['device'] = device
|
||
if batch_size is not None:
|
||
our_override_cfg['dataset'] = {'batch_size': batch_size}
|
||
if override_cfg is None:
|
||
override_cfg = {}
|
||
override_cfg = omegaconf.OmegaConf.merge(
|
||
omegaconf.DictConfig(override_cfg), omegaconf.DictConfig(our_override_cfg)) # type: ignore
|
||
solver = train.get_solver_from_sig(
|
||
sig, override_cfg=override_cfg,
|
||
load_best=True, disable_fsdp=True,
|
||
ignore_state_keys=['optimizer', 'ema'], **kwargs)
|
||
solver.model.eval()
|
||
return solver</code></pre>
|
||
</details>
|
||
<div class="desc"><p>Standard solver for AudioCraft.</p>
|
||
<p>The standard solver implements a base training loop with the following stages:
|
||
train, valid, evaluate and generate that are expected to be all defined for
|
||
solvers in AudioCraft. It also provides a nice default management of Dora history replay,
|
||
checkpoint management across epoch, and logging configuration.</p>
|
||
<p>AudioCraft solvers must inherit from the StandardSolver and define the methods
|
||
associated to each stage as well as the show, build_model and build_dataloaders methods.</p></div>
|
||
<h3>Ancestors</h3>
|
||
<ul class="hlist">
|
||
<li>abc.ABC</li>
|
||
<li>flashy.solver.BaseSolver</li>
|
||
</ul>
|
||
<h3>Subclasses</h3>
|
||
<ul class="hlist">
|
||
<li><a title="audiocraft.solvers.compression.CompressionSolver" href="compression.html#audiocraft.solvers.compression.CompressionSolver">CompressionSolver</a></li>
|
||
<li><a title="audiocraft.solvers.diffusion.DiffusionSolver" href="diffusion.html#audiocraft.solvers.diffusion.DiffusionSolver">DiffusionSolver</a></li>
|
||
<li><a title="audiocraft.solvers.musicgen.MusicGenSolver" href="musicgen.html#audiocraft.solvers.musicgen.MusicGenSolver">MusicGenSolver</a></li>
|
||
<li><a title="audiocraft.solvers.watermark.WatermarkSolver" href="watermark.html#audiocraft.solvers.watermark.WatermarkSolver">WatermarkSolver</a></li>
|
||
</ul>
|
||
<h3>Static methods</h3>
|
||
<dl>
|
||
<dt id="audiocraft.solvers.base.StandardSolver.get_eval_solver_from_sig"><code class="name flex">
|
||
<span>def <span class="ident">get_eval_solver_from_sig</span></span>(<span>sig: str,<br>dtype: str | None = None,<br>device: str | None = None,<br>autocast: bool = True,<br>batch_size: int | None = None,<br>override_cfg: dict | omegaconf.dictconfig.DictConfig | None = None,<br>**kwargs)</span>
|
||
</code></dt>
|
||
<dd>
|
||
<details class="source">
|
||
<summary>
|
||
<span>Expand source code</span>
|
||
</summary>
|
||
<pre><code class="python">@staticmethod
|
||
def get_eval_solver_from_sig(sig: str, dtype: tp.Optional[str] = None,
|
||
device: tp.Optional[str] = None, autocast: bool = True,
|
||
batch_size: tp.Optional[int] = None,
|
||
override_cfg: tp.Optional[tp.Union[dict, omegaconf.DictConfig]] = None,
|
||
**kwargs):
|
||
"""Mostly a convenience function around audiocraft.train.get_solver_from_sig,
|
||
populating all the proper param, deactivating EMA, FSDP, loading the best state,
|
||
basically all you need to get a solver ready to "play" with in single GPU mode
|
||
and with minimal memory overhead.
|
||
|
||
Args:
|
||
sig (str): signature to load.
|
||
dtype (str or None): potential dtype, as a string, i.e. 'float16'.
|
||
device (str or None): potential device, as a string, i.e. 'cuda'.
|
||
override_cfg (dict or omegaconf.DictConfig or None): potential device, as a string, i.e. 'cuda'.
|
||
"""
|
||
from audiocraft import train
|
||
our_override_cfg: tp.Dict[str, tp.Any] = {'optim': {'ema': {'use': False}}}
|
||
our_override_cfg['autocast'] = autocast
|
||
if dtype is not None:
|
||
our_override_cfg['dtype'] = dtype
|
||
if device is not None:
|
||
our_override_cfg['device'] = device
|
||
if batch_size is not None:
|
||
our_override_cfg['dataset'] = {'batch_size': batch_size}
|
||
if override_cfg is None:
|
||
override_cfg = {}
|
||
override_cfg = omegaconf.OmegaConf.merge(
|
||
omegaconf.DictConfig(override_cfg), omegaconf.DictConfig(our_override_cfg)) # type: ignore
|
||
solver = train.get_solver_from_sig(
|
||
sig, override_cfg=override_cfg,
|
||
load_best=True, disable_fsdp=True,
|
||
ignore_state_keys=['optimizer', 'ema'], **kwargs)
|
||
solver.model.eval()
|
||
return solver</code></pre>
|
||
</details>
|
||
<div class="desc"><p>Mostly a convenience function around audiocraft.train.get_solver_from_sig,
|
||
populating all the proper param, deactivating EMA, FSDP, loading the best state,
|
||
basically all you need to get a solver ready to "play" with in single GPU mode
|
||
and with minimal memory overhead.</p>
|
||
<h2 id="args">Args</h2>
|
||
<dl>
|
||
<dt><strong><code>sig</code></strong> : <code>str</code></dt>
|
||
<dd>signature to load.</dd>
|
||
<dt><strong><code>dtype</code></strong> : <code>str</code> or <code>None</code></dt>
|
||
<dd>potential dtype, as a string, i.e. 'float16'.</dd>
|
||
<dt><strong><code>device</code></strong> : <code>str</code> or <code>None</code></dt>
|
||
<dd>potential device, as a string, i.e. 'cuda'.</dd>
|
||
<dt><strong><code>override_cfg</code></strong> : <code>dict</code> or <code>omegaconf.DictConfig</code> or <code>None</code></dt>
|
||
<dd>potential device, as a string, i.e. 'cuda'.</dd>
|
||
</dl></div>
|
||
</dd>
|
||
</dl>
|
||
<h3>Instance variables</h3>
|
||
<dl>
|
||
<dt id="audiocraft.solvers.base.StandardSolver.autocast"><code class="name">prop <span class="ident">autocast</span></code></dt>
|
||
<dd>
|
||
<details class="source">
|
||
<summary>
|
||
<span>Expand source code</span>
|
||
</summary>
|
||
<pre><code class="python">@property
|
||
def autocast(self):
|
||
"""Convenient autocast (or not) using the solver configuration."""
|
||
return TorchAutocast(enabled=self.cfg.autocast, device_type=self.device, dtype=self.autocast_dtype)</code></pre>
|
||
</details>
|
||
<div class="desc"><p>Convenient autocast (or not) using the solver configuration.</p></div>
|
||
</dd>
|
||
<dt id="audiocraft.solvers.base.StandardSolver.best_metric_name"><code class="name">prop <span class="ident">best_metric_name</span> : str | None</code></dt>
|
||
<dd>
|
||
<details class="source">
|
||
<summary>
|
||
<span>Expand source code</span>
|
||
</summary>
|
||
<pre><code class="python">@property
|
||
def best_metric_name(self) -> tp.Optional[str]:
|
||
"""Metric name used to identify the best state. This metric should be stored in the metrics
|
||
used on the stage for best state identification (most likely, `valid`). If None, then
|
||
no best state is saved.
|
||
"""
|
||
return None</code></pre>
|
||
</details>
|
||
<div class="desc"><p>Metric name used to identify the best state. This metric should be stored in the metrics
|
||
used on the stage for best state identification (most likely, <code>valid</code>). If None, then
|
||
no best state is saved.</p></div>
|
||
</dd>
|
||
<dt id="audiocraft.solvers.base.StandardSolver.is_training"><code class="name">prop <span class="ident">is_training</span></code></dt>
|
||
<dd>
|
||
<details class="source">
|
||
<summary>
|
||
<span>Expand source code</span>
|
||
</summary>
|
||
<pre><code class="python">@property
|
||
def is_training(self):
|
||
return self.current_stage == 'train'</code></pre>
|
||
</details>
|
||
<div class="desc"></div>
|
||
</dd>
|
||
<dt id="audiocraft.solvers.base.StandardSolver.log_updates"><code class="name">prop <span class="ident">log_updates</span></code></dt>
|
||
<dd>
|
||
<details class="source">
|
||
<summary>
|
||
<span>Expand source code</span>
|
||
</summary>
|
||
<pre><code class="python">@property
|
||
def log_updates(self):
|
||
# convenient access to log updates
|
||
return self._log_updates</code></pre>
|
||
</details>
|
||
<div class="desc"></div>
|
||
</dd>
|
||
</dl>
|
||
<h3>Methods</h3>
|
||
<dl>
|
||
<dt id="audiocraft.solvers.base.StandardSolver.build_dataloaders"><code class="name flex">
|
||
<span>def <span class="ident">build_dataloaders</span></span>(<span>self)</span>
|
||
</code></dt>
|
||
<dd>
|
||
<details class="source">
|
||
<summary>
|
||
<span>Expand source code</span>
|
||
</summary>
|
||
<pre><code class="python">@abstractmethod
|
||
def build_dataloaders(self):
|
||
"""Method to implement to initialize dataloaders."""
|
||
...</code></pre>
|
||
</details>
|
||
<div class="desc"><p>Method to implement to initialize dataloaders.</p></div>
|
||
</dd>
|
||
<dt id="audiocraft.solvers.base.StandardSolver.build_model"><code class="name flex">
|
||
<span>def <span class="ident">build_model</span></span>(<span>self)</span>
|
||
</code></dt>
|
||
<dd>
|
||
<details class="source">
|
||
<summary>
|
||
<span>Expand source code</span>
|
||
</summary>
|
||
<pre><code class="python">@abstractmethod
|
||
def build_model(self):
|
||
"""Method to implement to initialize model."""
|
||
...</code></pre>
|
||
</details>
|
||
<div class="desc"><p>Method to implement to initialize model.</p></div>
|
||
</dd>
|
||
<dt id="audiocraft.solvers.base.StandardSolver.checkpoint_path"><code class="name flex">
|
||
<span>def <span class="ident">checkpoint_path</span></span>(<span>self, **kwargs)</span>
|
||
</code></dt>
|
||
<dd>
|
||
<details class="source">
|
||
<summary>
|
||
<span>Expand source code</span>
|
||
</summary>
|
||
<pre><code class="python">def checkpoint_path(self, **kwargs):
|
||
kwargs.setdefault('use_fsdp', self.cfg.fsdp.use)
|
||
return self.folder / checkpoint.checkpoint_name(**kwargs)</code></pre>
|
||
</details>
|
||
<div class="desc"></div>
|
||
</dd>
|
||
<dt id="audiocraft.solvers.base.StandardSolver.checkpoint_path_with_name"><code class="name flex">
|
||
<span>def <span class="ident">checkpoint_path_with_name</span></span>(<span>self, name: str, **kwargs)</span>
|
||
</code></dt>
|
||
<dd>
|
||
<details class="source">
|
||
<summary>
|
||
<span>Expand source code</span>
|
||
</summary>
|
||
<pre><code class="python">def checkpoint_path_with_name(self, name: str, **kwargs):
|
||
kwargs.setdefault('use_fsdp', self.cfg.fsdp.use)
|
||
return self.folder / checkpoint.checkpoint_name(name=name, **kwargs)</code></pre>
|
||
</details>
|
||
<div class="desc"></div>
|
||
</dd>
|
||
<dt id="audiocraft.solvers.base.StandardSolver.commit"><code class="name flex">
|
||
<span>def <span class="ident">commit</span></span>(<span>self, save_checkpoints: bool = True)</span>
|
||
</code></dt>
|
||
<dd>
|
||
<details class="source">
|
||
<summary>
|
||
<span>Expand source code</span>
|
||
</summary>
|
||
<pre><code class="python">def commit(self, save_checkpoints: bool = True):
|
||
"""Commit metrics to dora and save checkpoints at the end of an epoch."""
|
||
# we override commit to introduce more complex checkpoint saving behaviors
|
||
self.history.append(self._pending_metrics) # This will increase self.epoch
|
||
if save_checkpoints:
|
||
self.save_checkpoints()
|
||
self._start_epoch()
|
||
if flashy.distrib.is_rank_zero():
|
||
self.xp.link.update_history(self.history)</code></pre>
|
||
</details>
|
||
<div class="desc"><p>Commit metrics to dora and save checkpoints at the end of an epoch.</p></div>
|
||
</dd>
|
||
<dt id="audiocraft.solvers.base.StandardSolver.common_train_valid"><code class="name flex">
|
||
<span>def <span class="ident">common_train_valid</span></span>(<span>self, dataset_split: str, **kwargs: Any)</span>
|
||
</code></dt>
|
||
<dd>
|
||
<details class="source">
|
||
<summary>
|
||
<span>Expand source code</span>
|
||
</summary>
|
||
<pre><code class="python">def common_train_valid(self, dataset_split: str, **kwargs: tp.Any):
|
||
"""Common logic for train and valid stages."""
|
||
self.model.train(self.is_training)
|
||
|
||
loader = self.dataloaders[dataset_split]
|
||
# get a different order for distributed training, otherwise this will get ignored
|
||
if flashy.distrib.world_size() > 1 \
|
||
and isinstance(loader.sampler, torch.utils.data.distributed.DistributedSampler):
|
||
loader.sampler.set_epoch(self.epoch)
|
||
updates_per_epoch = self.train_updates_per_epoch if self.is_training else len(loader)
|
||
if self.cfg.benchmark_no_load:
|
||
self.logger.warning("Fake loading for benchmarking: re-using first batch")
|
||
batch = next(iter(loader))
|
||
loader = [batch] * updates_per_epoch # type: ignore
|
||
lp = self.log_progress(self.current_stage, loader, total=updates_per_epoch, updates=self.log_updates)
|
||
average = flashy.averager() # epoch wise average
|
||
instant_average = flashy.averager() # average between two logging
|
||
metrics: dict = {}
|
||
|
||
with self.profiler, self.deadlock_detect: # profiler will only run for the first 20 updates.
|
||
for idx, batch in enumerate(lp):
|
||
self.deadlock_detect.update('batch')
|
||
if idx >= updates_per_epoch:
|
||
break
|
||
metrics = {}
|
||
metrics = self.run_step(idx, batch, metrics)
|
||
self.deadlock_detect.update('step')
|
||
# run EMA step
|
||
if self.ema is not None and self.is_training and (idx + 1) % self.cfg.optim.ema.updates == 0:
|
||
self.logger.debug("EMA model step")
|
||
self.ema.step()
|
||
self.deadlock_detect.update('ema')
|
||
self.profiler.step()
|
||
instant_metrics = instant_average(metrics)
|
||
if lp.update(**instant_metrics):
|
||
instant_average = flashy.averager() # reset averager between two logging
|
||
metrics = average(metrics) # epoch wise average
|
||
self.deadlock_detect.update('end_batch')
|
||
|
||
metrics = flashy.distrib.average_metrics(metrics, updates_per_epoch)
|
||
return metrics</code></pre>
|
||
</details>
|
||
<div class="desc"><p>Common logic for train and valid stages.</p></div>
|
||
</dd>
|
||
<dt id="audiocraft.solvers.base.StandardSolver.epoch_checkpoint_path"><code class="name flex">
|
||
<span>def <span class="ident">epoch_checkpoint_path</span></span>(<span>self, epoch: int, **kwargs)</span>
|
||
</code></dt>
|
||
<dd>
|
||
<details class="source">
|
||
<summary>
|
||
<span>Expand source code</span>
|
||
</summary>
|
||
<pre><code class="python">def epoch_checkpoint_path(self, epoch: int, **kwargs):
|
||
kwargs.setdefault('use_fsdp', self.cfg.fsdp.use)
|
||
return self.folder / checkpoint.checkpoint_name(str(epoch), **kwargs)</code></pre>
|
||
</details>
|
||
<div class="desc"></div>
|
||
</dd>
|
||
<dt id="audiocraft.solvers.base.StandardSolver.evaluate"><code class="name flex">
|
||
<span>def <span class="ident">evaluate</span></span>(<span>self)</span>
|
||
</code></dt>
|
||
<dd>
|
||
<details class="source">
|
||
<summary>
|
||
<span>Expand source code</span>
|
||
</summary>
|
||
<pre><code class="python">@abstractmethod
|
||
def evaluate(self):
|
||
"""Evaluate stage."""
|
||
...</code></pre>
|
||
</details>
|
||
<div class="desc"><p>Evaluate stage.</p></div>
|
||
</dd>
|
||
<dt id="audiocraft.solvers.base.StandardSolver.generate"><code class="name flex">
|
||
<span>def <span class="ident">generate</span></span>(<span>self)</span>
|
||
</code></dt>
|
||
<dd>
|
||
<details class="source">
|
||
<summary>
|
||
<span>Expand source code</span>
|
||
</summary>
|
||
<pre><code class="python">@abstractmethod
|
||
def generate(self):
|
||
"""Generate stage."""
|
||
...</code></pre>
|
||
</details>
|
||
<div class="desc"><p>Generate stage.</p></div>
|
||
</dd>
|
||
<dt id="audiocraft.solvers.base.StandardSolver.initialize_ema"><code class="name flex">
|
||
<span>def <span class="ident">initialize_ema</span></span>(<span>self)</span>
|
||
</code></dt>
|
||
<dd>
|
||
<details class="source">
|
||
<summary>
|
||
<span>Expand source code</span>
|
||
</summary>
|
||
<pre><code class="python">def initialize_ema(self):
|
||
"""Initialize exponential moving average with the registered sources.
|
||
EMA object is created if the optim.ema.model.decay value is non-null.
|
||
"""
|
||
from .builders import get_ema
|
||
self.ema = get_ema(self._ema_sources, self.cfg.optim.ema)
|
||
if self.ema is None:
|
||
self.logger.info('No EMA on the model.')
|
||
else:
|
||
assert self.cfg.optim.ema.updates > 0
|
||
self.logger.info(
|
||
f'Initializing EMA on the model with decay = {self.ema.decay}'
|
||
f' every {self.cfg.optim.ema.updates} updates'
|
||
)</code></pre>
|
||
</details>
|
||
<div class="desc"><p>Initialize exponential moving average with the registered sources.
|
||
EMA object is created if the optim.ema.model.decay value is non-null.</p></div>
|
||
</dd>
|
||
<dt id="audiocraft.solvers.base.StandardSolver.load_checkpoints"><code class="name flex">
|
||
<span>def <span class="ident">load_checkpoints</span></span>(<span>self, load_best: bool = False, ignore_state_keys: List[str] = []) ‑> dict | None</span>
|
||
</code></dt>
|
||
<dd>
|
||
<details class="source">
|
||
<summary>
|
||
<span>Expand source code</span>
|
||
</summary>
|
||
<pre><code class="python">def load_checkpoints(self, load_best: bool = False, ignore_state_keys: tp.List[str] = []) -> tp.Optional[dict]:
|
||
"""Load last checkpoint or the one specified in continue_from.
|
||
|
||
Args:
|
||
load_best (bool): Whether to load from best state dict or not.
|
||
Best state dict is always used when not loading the current xp.
|
||
ignore_state_keys (list of str): List of sources to ignore when loading the state, e.g. `optimizer`.
|
||
Returns:
|
||
state (dict, optional): The loaded state dictionary.
|
||
"""
|
||
# load checkpoints from xp folder or cfg.continue_from
|
||
is_sharded = self.cfg.fsdp.use
|
||
load_from_path: tp.Optional[Path] = None
|
||
checkpoint_source: tp.Optional[checkpoint.CheckpointSource] = None
|
||
|
||
if load_best:
|
||
self.logger.info("Trying to load state_dict from best state.")
|
||
|
||
state: tp.Optional[dict] = None
|
||
rank0_checkpoint_path = self.checkpoint_path(use_fsdp=False)
|
||
current_checkpoint_path = self.checkpoint_path()
|
||
_pretrained_prefix = '//pretrained/'
|
||
continue_pretrained = (self.cfg.continue_from or '').startswith(_pretrained_prefix)
|
||
if rank0_checkpoint_path.exists():
|
||
self.logger.info(f"Loading existing checkpoint: {current_checkpoint_path}")
|
||
load_from_path = current_checkpoint_path
|
||
checkpoint.check_sharded_checkpoint(current_checkpoint_path, rank0_checkpoint_path)
|
||
checkpoint_source = checkpoint.CheckpointSource.CURRENT_XP
|
||
elif self.cfg.continue_from and not continue_pretrained:
|
||
self.logger.info(f"Continuing from provided checkpoint: {self.cfg.continue_from}")
|
||
# we're always continuing from consolidated checkpoints: self.cfg.use_fsdp and not continue_best
|
||
load_from_path = checkpoint.resolve_checkpoint_path(self.cfg.continue_from, use_fsdp=False)
|
||
if load_from_path is None:
|
||
self.logger.error('Could not resolve the continue_from checkpoint %s', self.cfg.continue_from)
|
||
raise RuntimeError(f'Could not resolve continue_from checkpoint {self.cfg.continue_from}')
|
||
checkpoint_source = checkpoint.CheckpointSource.OTHER
|
||
|
||
if load_from_path is not None:
|
||
state = checkpoint.load_checkpoint(load_from_path, is_sharded)
|
||
elif continue_pretrained:
|
||
self.logger.info("Loading a pretrained model. Ignoring 'load_best' and 'ignore_state_keys' params.")
|
||
state = self.load_from_pretrained(self.cfg.continue_from[len(_pretrained_prefix):])
|
||
checkpoint_source = checkpoint.CheckpointSource.PRETRAINED
|
||
load_best = True
|
||
|
||
# checkpoints are not from the current xp, we only retrieve the best state
|
||
if checkpoint_source is not None and checkpoint_source != checkpoint.CheckpointSource.CURRENT_XP:
|
||
assert state is not None
|
||
self.logger.info("Checkpoint source is not the current xp: Load state_dict from best state.")
|
||
load_best = True
|
||
state = {key: state[key] for key in self._continue_best_source_keys if key in state}
|
||
# loaded checkpoints are FSDP checkpoints: we're reading the best state
|
||
# from FSDP and we drop the regular best_state
|
||
if 'fsdp_best_state' in state and state['fsdp_best_state']:
|
||
state.pop('best_state', None)
|
||
self.logger.info("... Loaded checkpoint has FSDP best state")
|
||
# FSDP is enabled in the solver, if the loaded checkpoints do not have FSDP support
|
||
# then we're initializing FSDP best state with the regular best state
|
||
elif self.cfg.fsdp.use:
|
||
if 'fsdp_best_state' not in state or not state['fsdp_best_state']:
|
||
# we swap non-FSDP checkpoints best_state to FSDP-compatible best state
|
||
state['fsdp_best_state'] = state.pop('best_state')
|
||
self.logger.info("... Loaded checkpoint does not have FSDP best state. Use regular best state")
|
||
|
||
if state is not None:
|
||
if load_best:
|
||
self.logger.info("Ignoring keys when loading best %r", ignore_state_keys)
|
||
for key in set(ignore_state_keys):
|
||
if key in state:
|
||
state.pop(key)
|
||
has_best_state = 'best_state' in state or 'fsdp_best_state' in state
|
||
assert has_best_state, ("Trying to load best state but neither 'best_state'",
|
||
" or 'fsdp_best_state' found in checkpoints.")
|
||
self.load_state_dict(state)
|
||
|
||
# for FSDP, let's make extra sure nothing bad happened with out of sync
|
||
# checkpoints across workers.
|
||
epoch = float(self.epoch)
|
||
avg_epoch = flashy.distrib.average_metrics({'epoch': epoch})['epoch']
|
||
if avg_epoch != epoch:
|
||
raise RuntimeError(
|
||
f"Inconsistent loading of checkpoints happened, our epoch is {epoch} "
|
||
f"but average of epochs is {avg_epoch}, at least one gpu must have a "
|
||
"different epoch number.")
|
||
|
||
# on load_best, properly reinitialize state_dict, best states and ema
|
||
# otherwise we load from the current xp and don't alter anything
|
||
if load_best:
|
||
self.logger.info("Loading state_dict from best state.")
|
||
if not self.cfg.fsdp.use and self.fsdp_best_state:
|
||
# loading from an FSDP checkpoint but with FSDP deactivated
|
||
self.logger.info("... Loading from FSDP best state dict.")
|
||
self.best_state.load_state_dict(self.fsdp_best_state)
|
||
|
||
# if load_best, we permanently override the regular state_dict with the best state
|
||
if self.cfg.fsdp.use:
|
||
self.logger.info("FSDP is used, loading from FSDP best state.")
|
||
with fsdp.switch_to_full_state_dict(self._fsdp_modules):
|
||
# this might be really fragile but okay for now.
|
||
self.load_state_dict(self.fsdp_best_state)
|
||
else:
|
||
# we permanently swap the stateful objects to their best state
|
||
self._load_new_state_dict(self.best_state.state_dict())
|
||
|
||
# the EMA modules should also be instantiated with best state.
|
||
# the easiest way to do so is to reinitialize a new EMA with best state loaded.
|
||
if self.ema is not None:
|
||
self.logger.info("Re-initializing EMA from best state")
|
||
self.initialize_ema()
|
||
|
||
if self.cfg.fsdp.use:
|
||
self.logger.info("Re-initializing best state after using FSDP best state.")
|
||
for name in self.best_state.states.keys():
|
||
state_source = self._get_state_source(name)
|
||
self.best_state.update(name, state_source)
|
||
|
||
return state</code></pre>
|
||
</details>
|
||
<div class="desc"><p>Load last checkpoint or the one specified in continue_from.</p>
|
||
<h2 id="args">Args</h2>
|
||
<dl>
|
||
<dt><strong><code>load_best</code></strong> : <code>bool</code></dt>
|
||
<dd>Whether to load from best state dict or not.
|
||
Best state dict is always used when not loading the current xp.</dd>
|
||
<dt><strong><code>ignore_state_keys</code></strong> : <code>list</code> of <code>str</code></dt>
|
||
<dd>List of sources to ignore when loading the state, e.g. <code>optimizer</code>.</dd>
|
||
</dl>
|
||
<h2 id="returns">Returns</h2>
|
||
<p>state (dict, optional): The loaded state dictionary.</p></div>
|
||
</dd>
|
||
<dt id="audiocraft.solvers.base.StandardSolver.load_from_pretrained"><code class="name flex">
|
||
<span>def <span class="ident">load_from_pretrained</span></span>(<span>self, name: str) ‑> dict</span>
|
||
</code></dt>
|
||
<dd>
|
||
<details class="source">
|
||
<summary>
|
||
<span>Expand source code</span>
|
||
</summary>
|
||
<pre><code class="python">def load_from_pretrained(self, name: str) -> dict:
|
||
raise NotImplementedError("Solver does not provide a way to load pretrained models.")</code></pre>
|
||
</details>
|
||
<div class="desc"></div>
|
||
</dd>
|
||
<dt id="audiocraft.solvers.base.StandardSolver.log_model_summary"><code class="name flex">
|
||
<span>def <span class="ident">log_model_summary</span></span>(<span>self, model: torch.nn.modules.module.Module)</span>
|
||
</code></dt>
|
||
<dd>
|
||
<details class="source">
|
||
<summary>
|
||
<span>Expand source code</span>
|
||
</summary>
|
||
<pre><code class="python">def log_model_summary(self, model: nn.Module):
|
||
"""Log model summary, architecture and size of the model."""
|
||
self.logger.info(model)
|
||
mb = sum(p.numel() for p in model.parameters()) * 4 / 2 ** 20
|
||
self.logger.info("Size: %.1f MB", mb)</code></pre>
|
||
</details>
|
||
<div class="desc"><p>Log model summary, architecture and size of the model.</p></div>
|
||
</dd>
|
||
<dt id="audiocraft.solvers.base.StandardSolver.register_best_state"><code class="name flex">
|
||
<span>def <span class="ident">register_best_state</span></span>(<span>self, *args: str)</span>
|
||
</code></dt>
|
||
<dd>
|
||
<details class="source">
|
||
<summary>
|
||
<span>Expand source code</span>
|
||
</summary>
|
||
<pre><code class="python">def register_best_state(self, *args: str):
|
||
"""Register state sources in `BestStateDictManager` to keep their best states along with their
|
||
latest states. The best state will be used at evaluation stages instead of the latest states.
|
||
|
||
Shortcut around `BestStateDictManager.register` method. You can pass any number of
|
||
attribute, included nested attributes and those will be included into the checkpoints
|
||
and automatically restored when `BaseSolver.restore` is called.
|
||
"""
|
||
for name in args:
|
||
state_source = self._get_state_source(name)
|
||
assert name in self.stateful.sources, "Registered states in best should be registered in stateful first!"
|
||
self.best_state.register(name, state_source)</code></pre>
|
||
</details>
|
||
<div class="desc"><p>Register state sources in <code>BestStateDictManager</code> to keep their best states along with their
|
||
latest states. The best state will be used at evaluation stages instead of the latest states.</p>
|
||
<p>Shortcut around <code>BestStateDictManager.register</code> method. You can pass any number of
|
||
attribute, included nested attributes and those will be included into the checkpoints
|
||
and automatically restored when <code>BaseSolver.restore</code> is called.</p></div>
|
||
</dd>
|
||
<dt id="audiocraft.solvers.base.StandardSolver.register_ema"><code class="name flex">
|
||
<span>def <span class="ident">register_ema</span></span>(<span>self, *args: str)</span>
|
||
</code></dt>
|
||
<dd>
|
||
<details class="source">
|
||
<summary>
|
||
<span>Expand source code</span>
|
||
</summary>
|
||
<pre><code class="python">def register_ema(self, *args: str):
|
||
"""Register state sources for exponential moving average.
|
||
|
||
The registered sources are used to instantiate a ModuleDictEMA instance.
|
||
The ModuleDictEMA keeps a `nn.ModuleDict` module that is updated when self.ema.step() is called
|
||
and swapped with the original state sources with self.swap_ema_state() method.
|
||
|
||
Usage:
|
||
self.register_ema('model')
|
||
"""
|
||
assert self.ema is None, "Cannot register state source to already instantiated EMA."
|
||
for name in args:
|
||
self._ema_sources[name] = getattr(self, name)</code></pre>
|
||
</details>
|
||
<div class="desc"><p>Register state sources for exponential moving average.</p>
|
||
<p>The registered sources are used to instantiate a ModuleDictEMA instance.
|
||
The ModuleDictEMA keeps a <code>nn.ModuleDict</code> module that is updated when self.ema.step() is called
|
||
and swapped with the original state sources with self.swap_ema_state() method.</p>
|
||
<h2 id="usage">Usage</h2>
|
||
<p>self.register_ema('model')</p></div>
|
||
</dd>
|
||
<dt id="audiocraft.solvers.base.StandardSolver.restore"><code class="name flex">
|
||
<span>def <span class="ident">restore</span></span>(<span>self,<br>load_best: bool = False,<br>replay_metrics: bool = False,<br>ignore_state_keys: List[str] = []) ‑> bool</span>
|
||
</code></dt>
|
||
<dd>
|
||
<details class="source">
|
||
<summary>
|
||
<span>Expand source code</span>
|
||
</summary>
|
||
<pre><code class="python">def restore(self, load_best: bool = False, replay_metrics: bool = False,
|
||
ignore_state_keys: tp.List[str] = []) -> bool:
|
||
"""Restore the status of a solver for a given xp.
|
||
|
||
Args:
|
||
load_best (bool): if `True`, load the best state from the checkpoint.
|
||
replay_metrics (bool): if `True`, logs all the metrics from past epochs.
|
||
ignore_state_keys (list of str): list of sources to ignore when loading the state, e.g. `optimizer`.
|
||
"""
|
||
self.logger.info("Restoring weights and history.")
|
||
restored_checkpoints = self.load_checkpoints(load_best, ignore_state_keys)
|
||
|
||
self.logger.info("Model hash: %s", model_hash(self.model))
|
||
|
||
if replay_metrics and len(self.history) > 0:
|
||
self.logger.info("Replaying past metrics...")
|
||
for epoch, stages in enumerate(self.history):
|
||
for stage_name, metrics in stages.items():
|
||
# We manually log the metrics summary to the result logger
|
||
# as we don't want to add them to the pending metrics
|
||
self.result_logger._log_summary(stage_name, metrics, step=epoch + 1, step_name='epoch',
|
||
formatter=self.get_formatter(stage_name))
|
||
return restored_checkpoints is not None</code></pre>
|
||
</details>
|
||
<div class="desc"><p>Restore the status of a solver for a given xp.</p>
|
||
<h2 id="args">Args</h2>
|
||
<dl>
|
||
<dt><strong><code>load_best</code></strong> : <code>bool</code></dt>
|
||
<dd>if <code>True</code>, load the best state from the checkpoint.</dd>
|
||
<dt><strong><code>replay_metrics</code></strong> : <code>bool</code></dt>
|
||
<dd>if <code>True</code>, logs all the metrics from past epochs.</dd>
|
||
<dt><strong><code>ignore_state_keys</code></strong> : <code>list</code> of <code>str</code></dt>
|
||
<dd>list of sources to ignore when loading the state, e.g. <code>optimizer</code>.</dd>
|
||
</dl></div>
|
||
</dd>
|
||
<dt id="audiocraft.solvers.base.StandardSolver.run"><code class="name flex">
|
||
<span>def <span class="ident">run</span></span>(<span>self)</span>
|
||
</code></dt>
|
||
<dd>
|
||
<details class="source">
|
||
<summary>
|
||
<span>Expand source code</span>
|
||
</summary>
|
||
<pre><code class="python">def run(self):
|
||
"""Training loop."""
|
||
assert len(self.state_dict()) > 0
|
||
self.restore(replay_metrics=True) # load checkpoint and replay history
|
||
self.log_hyperparams(dict_from_config(self.cfg))
|
||
for epoch in range(self.epoch, self.cfg.optim.epochs + 1):
|
||
if self.should_stop_training():
|
||
return
|
||
self.run_epoch()
|
||
# Commit will send the metrics to Dora and save checkpoints by default.
|
||
self.commit()</code></pre>
|
||
</details>
|
||
<div class="desc"><p>Training loop.</p></div>
|
||
</dd>
|
||
<dt id="audiocraft.solvers.base.StandardSolver.run_epoch"><code class="name flex">
|
||
<span>def <span class="ident">run_epoch</span></span>(<span>self)</span>
|
||
</code></dt>
|
||
<dd>
|
||
<details class="source">
|
||
<summary>
|
||
<span>Expand source code</span>
|
||
</summary>
|
||
<pre><code class="python">def run_epoch(self):
|
||
"""Run a single epoch with all stages.
|
||
|
||
Metrics for a given stage are stored in _pending_metrics and committed by the solver afterwards.
|
||
Children solvers can extend this method with custom behavior, e.g.:
|
||
|
||
def run_epoch(self):
|
||
... # custom code
|
||
super().run_epoch()
|
||
... # custom code
|
||
"""
|
||
self.run_stage('train', self.train)
|
||
with torch.no_grad():
|
||
with self.swap_ema_state():
|
||
self.run_stage('valid', self.valid)
|
||
# the best state is updated with EMA states if available
|
||
self.update_best_state_from_stage('valid')
|
||
with self.swap_best_state():
|
||
if self.should_run_stage('evaluate'):
|
||
self.run_stage('evaluate', self.evaluate)
|
||
if self.should_run_stage('generate'):
|
||
self.run_stage('generate', with_rank_rng()(self.generate))</code></pre>
|
||
</details>
|
||
<div class="desc"><p>Run a single epoch with all stages.</p>
|
||
<p>Metrics for a given stage are stored in _pending_metrics and committed by the solver afterwards.
|
||
Children solvers can extend this method with custom behavior, e.g.:</p>
|
||
<pre><code>def run_epoch(self):
|
||
... # custom code
|
||
super().run_epoch()
|
||
... # custom code
|
||
</code></pre></div>
|
||
</dd>
|
||
<dt id="audiocraft.solvers.base.StandardSolver.run_one_stage"><code class="name flex">
|
||
<span>def <span class="ident">run_one_stage</span></span>(<span>self, stage_name: str)</span>
|
||
</code></dt>
|
||
<dd>
|
||
<details class="source">
|
||
<summary>
|
||
<span>Expand source code</span>
|
||
</summary>
|
||
<pre><code class="python">def run_one_stage(self, stage_name: str):
|
||
"""Run only the specified stage.
|
||
This method is useful to only generate samples from a trained experiment
|
||
or rerun the validation or evaluation stages.
|
||
"""
|
||
fn = {
|
||
'generate': with_rank_rng()(self.generate),
|
||
'evaluate': self.evaluate,
|
||
'valid': self.valid,
|
||
}
|
||
if stage_name not in fn:
|
||
raise ValueError(f'Trying to run stage {stage_name} is not supported.')
|
||
assert len(self.state_dict()) > 0
|
||
self._start_epoch()
|
||
with torch.no_grad(), self.swap_best_state():
|
||
self.run_stage(stage_name, fn[stage_name])
|
||
if not self.cfg.execute_inplace:
|
||
self.commit(save_checkpoints=False)</code></pre>
|
||
</details>
|
||
<div class="desc"><p>Run only the specified stage.
|
||
This method is useful to only generate samples from a trained experiment
|
||
or rerun the validation or evaluation stages.</p></div>
|
||
</dd>
|
||
<dt id="audiocraft.solvers.base.StandardSolver.run_step"><code class="name flex">
|
||
<span>def <span class="ident">run_step</span></span>(<span>self, idx: int, batch: Any, metrics: dict)</span>
|
||
</code></dt>
|
||
<dd>
|
||
<details class="source">
|
||
<summary>
|
||
<span>Expand source code</span>
|
||
</summary>
|
||
<pre><code class="python">@abstractmethod
|
||
def run_step(self, idx: int, batch: tp.Any, metrics: dict):
|
||
"""Perform one training or valid step on a given batch."""
|
||
...</code></pre>
|
||
</details>
|
||
<div class="desc"><p>Perform one training or valid step on a given batch.</p></div>
|
||
</dd>
|
||
<dt id="audiocraft.solvers.base.StandardSolver.save_checkpoints"><code class="name flex">
|
||
<span>def <span class="ident">save_checkpoints</span></span>(<span>self)</span>
|
||
</code></dt>
|
||
<dd>
|
||
<details class="source">
|
||
<summary>
|
||
<span>Expand source code</span>
|
||
</summary>
|
||
<pre><code class="python">def save_checkpoints(self):
|
||
"""Save checkpoint, optionally keeping a copy for a given epoch."""
|
||
is_sharded = self.cfg.fsdp.use
|
||
if not flashy.distrib.is_rank_zero() and not is_sharded:
|
||
return
|
||
self.logger.info("Model hash: %s", model_hash(self.model))
|
||
state = self.state_dict()
|
||
epoch = self.epoch - 1 # pushing metrics will increase the epoch in Flashy, so we do -1 here
|
||
|
||
# save minimal state_dict as new checkpoint every X epoch
|
||
if self.cfg.checkpoint.save_every:
|
||
if epoch % self.cfg.checkpoint.save_every == 0:
|
||
minimal_state = state
|
||
if self.cfg.checkpoint.keep_every_states is not None and len(self.cfg.checkpoint.keep_every_states) > 0:
|
||
minimal_state = {
|
||
name: source for name, source in state.items()
|
||
if name in self.cfg.checkpoint.keep_every_states
|
||
}
|
||
epoch_checkpoint_path = self.epoch_checkpoint_path(epoch)
|
||
checkpoint.save_checkpoint(minimal_state, epoch_checkpoint_path, is_sharded)
|
||
|
||
# save checkpoint as latest checkpoint
|
||
if self.cfg.checkpoint.save_last:
|
||
last_checkpoint_path = self.checkpoint_path()
|
||
checkpoint.save_checkpoint(state, last_checkpoint_path, is_sharded)
|
||
|
||
# flush any stale checkpoint to reduce disk footprint
|
||
checkpoint.flush_stale_checkpoints(self.checkpoint_path())</code></pre>
|
||
</details>
|
||
<div class="desc"><p>Save checkpoint, optionally keeping a copy for a given epoch.</p></div>
|
||
</dd>
|
||
<dt id="audiocraft.solvers.base.StandardSolver.should_run_stage"><code class="name flex">
|
||
<span>def <span class="ident">should_run_stage</span></span>(<span>self, stage_name) ‑> bool</span>
|
||
</code></dt>
|
||
<dd>
|
||
<details class="source">
|
||
<summary>
|
||
<span>Expand source code</span>
|
||
</summary>
|
||
<pre><code class="python">def should_run_stage(self, stage_name) -> bool:
|
||
"""Check whether we want to run the specified stages."""
|
||
stage_every = self.cfg[stage_name].get('every', None)
|
||
is_last_epoch = self.epoch == self.cfg.optim.epochs
|
||
is_epoch_every = (stage_every and self.epoch % stage_every == 0)
|
||
return is_last_epoch or is_epoch_every</code></pre>
|
||
</details>
|
||
<div class="desc"><p>Check whether we want to run the specified stages.</p></div>
|
||
</dd>
|
||
<dt id="audiocraft.solvers.base.StandardSolver.should_stop_training"><code class="name flex">
|
||
<span>def <span class="ident">should_stop_training</span></span>(<span>self) ‑> bool</span>
|
||
</code></dt>
|
||
<dd>
|
||
<details class="source">
|
||
<summary>
|
||
<span>Expand source code</span>
|
||
</summary>
|
||
<pre><code class="python">def should_stop_training(self) -> bool:
|
||
"""Check whether we should stop training or not."""
|
||
return self.epoch > self.cfg.optim.epochs</code></pre>
|
||
</details>
|
||
<div class="desc"><p>Check whether we should stop training or not.</p></div>
|
||
</dd>
|
||
<dt id="audiocraft.solvers.base.StandardSolver.show"><code class="name flex">
|
||
<span>def <span class="ident">show</span></span>(<span>self)</span>
|
||
</code></dt>
|
||
<dd>
|
||
<details class="source">
|
||
<summary>
|
||
<span>Expand source code</span>
|
||
</summary>
|
||
<pre><code class="python">@abstractmethod
|
||
def show(self):
|
||
"""Method to log any information without running the job."""
|
||
...</code></pre>
|
||
</details>
|
||
<div class="desc"><p>Method to log any information without running the job.</p></div>
|
||
</dd>
|
||
<dt id="audiocraft.solvers.base.StandardSolver.swap_best_state"><code class="name flex">
|
||
<span>def <span class="ident">swap_best_state</span></span>(<span>self)</span>
|
||
</code></dt>
|
||
<dd>
|
||
<details class="source">
|
||
<summary>
|
||
<span>Expand source code</span>
|
||
</summary>
|
||
<pre><code class="python">@contextmanager
|
||
def swap_best_state(self):
|
||
self.logger.debug(f"Swapping to best state for: {', '.join(self.best_state.state_dict().keys())}")
|
||
old_states = self._load_new_state_dict(self.best_state.state_dict())
|
||
try:
|
||
yield
|
||
finally:
|
||
self.logger.debug("Swapping back from best to original state")
|
||
for name, old_state in old_states.items():
|
||
state_source = self._get_state_source(name)
|
||
state_source.load_state_dict(old_state)</code></pre>
|
||
</details>
|
||
<div class="desc"></div>
|
||
</dd>
|
||
<dt id="audiocraft.solvers.base.StandardSolver.swap_ema_state"><code class="name flex">
|
||
<span>def <span class="ident">swap_ema_state</span></span>(<span>self)</span>
|
||
</code></dt>
|
||
<dd>
|
||
<details class="source">
|
||
<summary>
|
||
<span>Expand source code</span>
|
||
</summary>
|
||
<pre><code class="python">@contextmanager
|
||
def swap_ema_state(self):
|
||
if self.ema is None:
|
||
yield
|
||
else:
|
||
ema_state_dict = self.ema.state_dict()['state']
|
||
self.logger.debug(f"Swapping to EMA state for: {', '.join(ema_state_dict.keys())}")
|
||
old_states = self._load_new_state_dict(ema_state_dict)
|
||
try:
|
||
yield
|
||
finally:
|
||
self.logger.debug("Swapping back from EMA state to original state")
|
||
for name, old_state in old_states.items():
|
||
state_source = self._get_state_source(name)
|
||
state_source.load_state_dict(old_state)</code></pre>
|
||
</details>
|
||
<div class="desc"></div>
|
||
</dd>
|
||
<dt id="audiocraft.solvers.base.StandardSolver.train"><code class="name flex">
|
||
<span>def <span class="ident">train</span></span>(<span>self)</span>
|
||
</code></dt>
|
||
<dd>
|
||
<details class="source">
|
||
<summary>
|
||
<span>Expand source code</span>
|
||
</summary>
|
||
<pre><code class="python">def train(self):
|
||
"""Train stage."""
|
||
return self.common_train_valid('train')</code></pre>
|
||
</details>
|
||
<div class="desc"><p>Train stage.</p></div>
|
||
</dd>
|
||
<dt id="audiocraft.solvers.base.StandardSolver.update_best_state_from_stage"><code class="name flex">
|
||
<span>def <span class="ident">update_best_state_from_stage</span></span>(<span>self, stage_name: str = 'valid')</span>
|
||
</code></dt>
|
||
<dd>
|
||
<details class="source">
|
||
<summary>
|
||
<span>Expand source code</span>
|
||
</summary>
|
||
<pre><code class="python">def update_best_state_from_stage(self, stage_name: str = 'valid'):
|
||
"""Update latest best state based on pending metrics of a given stage. This method relies
|
||
on the `BestStateDictManager.update` method to update the best state_dict with latest weights
|
||
if the registered states happen to match to the best performing setup.
|
||
"""
|
||
if self.best_metric_name is None:
|
||
# when no best metric is defined, the last state is always the best
|
||
self._new_best_state = True
|
||
self.logger.info("Updating best state with current state.")
|
||
else:
|
||
assert stage_name in self._pending_metrics, f"Metrics for stage {stage_name} not found."
|
||
assert self.best_metric_name in self._pending_metrics[stage_name], \
|
||
f"Best metric not found in {stage_name} metrics. Cannot register best state"
|
||
current_score = self._pending_metrics[stage_name][self.best_metric_name]
|
||
all_best_metric_scores = [
|
||
past_metrics[stage_name][self.best_metric_name]
|
||
for past_metrics in self.history
|
||
]
|
||
all_best_metric_scores.append(current_score)
|
||
best_score = min(all_best_metric_scores)
|
||
self._new_best_state = current_score == best_score
|
||
if self._new_best_state:
|
||
old_best = min(all_best_metric_scores[:-1] + [float('inf')])
|
||
self.logger.info(
|
||
f"New best state with {self.best_metric_name}={current_score:.3f} (was {old_best:.3f})")
|
||
|
||
if self._new_best_state:
|
||
if self.cfg.fsdp.use:
|
||
# this will give an empty state dict on all ranks but the rank 0
|
||
# which will have a copy in memory of the full model.
|
||
with fsdp.switch_to_full_state_dict(self._fsdp_modules):
|
||
for name in self.best_state.states.keys():
|
||
state_source = self._get_state_source(name)
|
||
self.best_state.update(name, state_source)
|
||
# we save to a different dict.
|
||
self.fsdp_best_state.update(self.best_state.state_dict())
|
||
# We cannot efficiently load fsdp_best_state when using FSDP,
|
||
# so we have do do a second pass, with the local shards.
|
||
for name in self.best_state.states.keys():
|
||
state_source = self._get_state_source(name)
|
||
self.best_state.update(name, state_source)</code></pre>
|
||
</details>
|
||
<div class="desc"><p>Update latest best state based on pending metrics of a given stage. This method relies
|
||
on the <code>BestStateDictManager.update</code> method to update the best state_dict with latest weights
|
||
if the registered states happen to match to the best performing setup.</p></div>
|
||
</dd>
|
||
<dt id="audiocraft.solvers.base.StandardSolver.valid"><code class="name flex">
|
||
<span>def <span class="ident">valid</span></span>(<span>self)</span>
|
||
</code></dt>
|
||
<dd>
|
||
<details class="source">
|
||
<summary>
|
||
<span>Expand source code</span>
|
||
</summary>
|
||
<pre><code class="python">def valid(self):
|
||
"""Valid stage."""
|
||
return self.common_train_valid('valid')</code></pre>
|
||
</details>
|
||
<div class="desc"><p>Valid stage.</p></div>
|
||
</dd>
|
||
<dt id="audiocraft.solvers.base.StandardSolver.wrap_with_fsdp"><code class="name flex">
|
||
<span>def <span class="ident">wrap_with_fsdp</span></span>(<span>self, model: torch.nn.modules.module.Module, *args, **kwargs)</span>
|
||
</code></dt>
|
||
<dd>
|
||
<details class="source">
|
||
<summary>
|
||
<span>Expand source code</span>
|
||
</summary>
|
||
<pre><code class="python">def wrap_with_fsdp(self, model: torch.nn.Module, *args, **kwargs):
|
||
model = fsdp.wrap_with_fsdp(self.cfg.fsdp, model, *args, **kwargs)
|
||
if isinstance(model, fsdp.FSDP):
|
||
self._fsdp_modules.append(model)
|
||
return model</code></pre>
|
||
</details>
|
||
<div class="desc"></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.solvers" href="index.html">audiocraft.solvers</a></code></li>
|
||
</ul>
|
||
</li>
|
||
<li><h3><a href="#header-classes">Classes</a></h3>
|
||
<ul>
|
||
<li>
|
||
<h4><code><a title="audiocraft.solvers.base.StandardSolver" href="#audiocraft.solvers.base.StandardSolver">StandardSolver</a></code></h4>
|
||
<ul class="">
|
||
<li><code><a title="audiocraft.solvers.base.StandardSolver.autocast" href="#audiocraft.solvers.base.StandardSolver.autocast">autocast</a></code></li>
|
||
<li><code><a title="audiocraft.solvers.base.StandardSolver.best_metric_name" href="#audiocraft.solvers.base.StandardSolver.best_metric_name">best_metric_name</a></code></li>
|
||
<li><code><a title="audiocraft.solvers.base.StandardSolver.build_dataloaders" href="#audiocraft.solvers.base.StandardSolver.build_dataloaders">build_dataloaders</a></code></li>
|
||
<li><code><a title="audiocraft.solvers.base.StandardSolver.build_model" href="#audiocraft.solvers.base.StandardSolver.build_model">build_model</a></code></li>
|
||
<li><code><a title="audiocraft.solvers.base.StandardSolver.checkpoint_path" href="#audiocraft.solvers.base.StandardSolver.checkpoint_path">checkpoint_path</a></code></li>
|
||
<li><code><a title="audiocraft.solvers.base.StandardSolver.checkpoint_path_with_name" href="#audiocraft.solvers.base.StandardSolver.checkpoint_path_with_name">checkpoint_path_with_name</a></code></li>
|
||
<li><code><a title="audiocraft.solvers.base.StandardSolver.commit" href="#audiocraft.solvers.base.StandardSolver.commit">commit</a></code></li>
|
||
<li><code><a title="audiocraft.solvers.base.StandardSolver.common_train_valid" href="#audiocraft.solvers.base.StandardSolver.common_train_valid">common_train_valid</a></code></li>
|
||
<li><code><a title="audiocraft.solvers.base.StandardSolver.epoch_checkpoint_path" href="#audiocraft.solvers.base.StandardSolver.epoch_checkpoint_path">epoch_checkpoint_path</a></code></li>
|
||
<li><code><a title="audiocraft.solvers.base.StandardSolver.evaluate" href="#audiocraft.solvers.base.StandardSolver.evaluate">evaluate</a></code></li>
|
||
<li><code><a title="audiocraft.solvers.base.StandardSolver.generate" href="#audiocraft.solvers.base.StandardSolver.generate">generate</a></code></li>
|
||
<li><code><a title="audiocraft.solvers.base.StandardSolver.get_eval_solver_from_sig" href="#audiocraft.solvers.base.StandardSolver.get_eval_solver_from_sig">get_eval_solver_from_sig</a></code></li>
|
||
<li><code><a title="audiocraft.solvers.base.StandardSolver.initialize_ema" href="#audiocraft.solvers.base.StandardSolver.initialize_ema">initialize_ema</a></code></li>
|
||
<li><code><a title="audiocraft.solvers.base.StandardSolver.is_training" href="#audiocraft.solvers.base.StandardSolver.is_training">is_training</a></code></li>
|
||
<li><code><a title="audiocraft.solvers.base.StandardSolver.load_checkpoints" href="#audiocraft.solvers.base.StandardSolver.load_checkpoints">load_checkpoints</a></code></li>
|
||
<li><code><a title="audiocraft.solvers.base.StandardSolver.load_from_pretrained" href="#audiocraft.solvers.base.StandardSolver.load_from_pretrained">load_from_pretrained</a></code></li>
|
||
<li><code><a title="audiocraft.solvers.base.StandardSolver.log_model_summary" href="#audiocraft.solvers.base.StandardSolver.log_model_summary">log_model_summary</a></code></li>
|
||
<li><code><a title="audiocraft.solvers.base.StandardSolver.log_updates" href="#audiocraft.solvers.base.StandardSolver.log_updates">log_updates</a></code></li>
|
||
<li><code><a title="audiocraft.solvers.base.StandardSolver.register_best_state" href="#audiocraft.solvers.base.StandardSolver.register_best_state">register_best_state</a></code></li>
|
||
<li><code><a title="audiocraft.solvers.base.StandardSolver.register_ema" href="#audiocraft.solvers.base.StandardSolver.register_ema">register_ema</a></code></li>
|
||
<li><code><a title="audiocraft.solvers.base.StandardSolver.restore" href="#audiocraft.solvers.base.StandardSolver.restore">restore</a></code></li>
|
||
<li><code><a title="audiocraft.solvers.base.StandardSolver.run" href="#audiocraft.solvers.base.StandardSolver.run">run</a></code></li>
|
||
<li><code><a title="audiocraft.solvers.base.StandardSolver.run_epoch" href="#audiocraft.solvers.base.StandardSolver.run_epoch">run_epoch</a></code></li>
|
||
<li><code><a title="audiocraft.solvers.base.StandardSolver.run_one_stage" href="#audiocraft.solvers.base.StandardSolver.run_one_stage">run_one_stage</a></code></li>
|
||
<li><code><a title="audiocraft.solvers.base.StandardSolver.run_step" href="#audiocraft.solvers.base.StandardSolver.run_step">run_step</a></code></li>
|
||
<li><code><a title="audiocraft.solvers.base.StandardSolver.save_checkpoints" href="#audiocraft.solvers.base.StandardSolver.save_checkpoints">save_checkpoints</a></code></li>
|
||
<li><code><a title="audiocraft.solvers.base.StandardSolver.should_run_stage" href="#audiocraft.solvers.base.StandardSolver.should_run_stage">should_run_stage</a></code></li>
|
||
<li><code><a title="audiocraft.solvers.base.StandardSolver.should_stop_training" href="#audiocraft.solvers.base.StandardSolver.should_stop_training">should_stop_training</a></code></li>
|
||
<li><code><a title="audiocraft.solvers.base.StandardSolver.show" href="#audiocraft.solvers.base.StandardSolver.show">show</a></code></li>
|
||
<li><code><a title="audiocraft.solvers.base.StandardSolver.swap_best_state" href="#audiocraft.solvers.base.StandardSolver.swap_best_state">swap_best_state</a></code></li>
|
||
<li><code><a title="audiocraft.solvers.base.StandardSolver.swap_ema_state" href="#audiocraft.solvers.base.StandardSolver.swap_ema_state">swap_ema_state</a></code></li>
|
||
<li><code><a title="audiocraft.solvers.base.StandardSolver.train" href="#audiocraft.solvers.base.StandardSolver.train">train</a></code></li>
|
||
<li><code><a title="audiocraft.solvers.base.StandardSolver.update_best_state_from_stage" href="#audiocraft.solvers.base.StandardSolver.update_best_state_from_stage">update_best_state_from_stage</a></code></li>
|
||
<li><code><a title="audiocraft.solvers.base.StandardSolver.valid" href="#audiocraft.solvers.base.StandardSolver.valid">valid</a></code></li>
|
||
<li><code><a title="audiocraft.solvers.base.StandardSolver.wrap_with_fsdp" href="#audiocraft.solvers.base.StandardSolver.wrap_with_fsdp">wrap_with_fsdp</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>
|