projects
/
mygptrnn.git
/ blobdiff
commit
grep
author
committer
pickaxe
?
search:
re
summary
|
shortlog
|
log
|
commit
|
commitdiff
|
tree
raw
|
inline
| side by side
Update.
[mygptrnn.git]
/
tasks.py
diff --git
a/tasks.py
b/tasks.py
index
afad8af
..
4777a11
100755
(executable)
--- a/
tasks.py
+++ b/
tasks.py
@@
-1515,11
+1515,13
@@
class Grid(Task):
self.train_input = self.str2tensor(self.train_descr)
self.test_input = self.str2tensor(self.test_descr)
self.train_input = self.str2tensor(self.train_descr)
self.test_input = self.str2tensor(self.test_descr)
- def batches(self, split="train"):
+ def batches(self, split="train"
, desc=None
):
assert split in {"train", "test"}
input = self.train_input if split == "train" else self.test_input
assert split in {"train", "test"}
input = self.train_input if split == "train" else self.test_input
+ if desc is None:
+ desc = f"epoch-{split}"
for batch in tqdm.tqdm(
for batch in tqdm.tqdm(
- input.split(self.batch_size), dynamic_ncols=True, desc=
f"epoch-{split}"
+ input.split(self.batch_size), dynamic_ncols=True, desc=
desc
):
yield self.trim(batch)
):
yield self.trim(batch)
@@
-1618,11
+1620,13
@@
class QMLP(Task):
self.nb_codes = max(self.train_input.max(), self.test_input.max()) + 1
self.nb_codes = max(self.train_input.max(), self.test_input.max()) + 1
- def batches(self, split="train"):
+ def batches(self, split="train"
, desc=None
):
assert split in {"train", "test"}
input = self.train_input if split == "train" else self.test_input
assert split in {"train", "test"}
input = self.train_input if split == "train" else self.test_input
+ if desc is None:
+ desc = f"epoch-{split}"
for batch in tqdm.tqdm(
for batch in tqdm.tqdm(
- input.split(self.batch_size), dynamic_ncols=True, desc=
f"epoch-{split}"
+ input.split(self.batch_size), dynamic_ncols=True, desc=
desc
):
yield batch
):
yield batch