mirror of
https://github.com/wassname/pytorch-lightning.git
synced 2026-10-04 13:00:36 +08:00
* remove the need for hparams * remove the need for hparams * remove the need for hparams * remove the need for hparams * replace self.hparams * replace self.hparams * replace self.hparams * replace self.hparams * replace self.hparams * replace self.hparams * replace self.hparams * replace self.hparams * replace self.hparams * replace self.hparams * replace self.hparams * replace self.hparams * replace self.hparams * replace self.hparams * replace self.hparams * replace self.hparams * replace self.hparams * replace self.hparams * replace self.hparams * replace self.hparams * replace self.hparams * replace self.hparams * replace self.hparams * replace self.hparams * replace self.hparams * replace self.hparams * replace self.hparams * replace self.hparams * replace self.hparams * replace self.hparams * replace self.hparams * replace self.hparams * fixed * fixed * fixed * fixed * fixed * fixed * fixed * fixed * fixed * fixed * fixed * fixed * fixed * fixed * finished moco * basic * testing * todo * recurse * hparams * persist * hparams * chlog * tests * tests * tests * tests * tests * tests * review * saving * tests * tests * tests * docs * finished moco * hparams * review * Apply suggestions from code review Co-authored-by: Adrian Wälchli <aedu.waelchli@gmail.com> * hparams * overwrite * transform * transform * transform * transform * cleaning * cleaning * tests * examples * examples * examples * Apply suggestions from code review Co-authored-by: Adrian Wälchli <aedu.waelchli@gmail.com> * chp key * tests * Apply suggestions from code review * class * updated docs * updated docs * updated docs * updated docs * save * wip * fix * flake8 Co-authored-by: Jirka <jirka@pytorchlightning.ai> Co-authored-by: Jirka Borovec <Borda@users.noreply.github.com> Co-authored-by: Adrian Wälchli <aedu.waelchli@gmail.com>
29 lines
747 B
Python
29 lines
747 B
Python
from torch.utils.data import DataLoader
|
|
|
|
from tests.base.datasets import TrialMNIST
|
|
|
|
|
|
class ModelTemplateData:
|
|
hparams: ...
|
|
|
|
def dataloader(self, train):
|
|
dataset = TrialMNIST(root=self.data_root, train=train, download=True)
|
|
|
|
loader = DataLoader(
|
|
dataset=dataset,
|
|
batch_size=self.batch_size,
|
|
# test and valid shall not be shuffled
|
|
shuffle=train,
|
|
)
|
|
return loader
|
|
|
|
|
|
class ModelTemplateUtils:
|
|
|
|
def get_output_metric(self, output, name):
|
|
if isinstance(output, dict):
|
|
val = output[name]
|
|
else: # if it is 2level deep -> per dataloader and per batch
|
|
val = sum(out[name] for out in output) / len(output)
|
|
return val
|