X-Git-Url: https://fleuret.org/cgi-bin/gitweb/gitweb.cgi?a=blobdiff_plain;f=tasks.py;h=0a4dd6fa2f880e93aecbbd494621fae26b7dcdbb;hb=a291e213a152364b74e833200191c08a36451a90;hp=e14ceb76e18ba38283acb99a5f214d42b92cc1b1;hpb=5703df4c32a0856c8fa4b1ff97810cdc1fb76253;p=picoclvr.git diff --git a/tasks.py b/tasks.py index e14ceb7..0a4dd6f 100755 --- a/tasks.py +++ b/tasks.py @@ -1042,7 +1042,7 @@ class RPL(Task): ) ], 0, - ).to(self.device) + ) def seq2str(self, seq): return " ".join([self.id2token[i] for i in seq]) @@ -1056,6 +1056,7 @@ class RPL(Task): max_input=9, prog_len=6, nb_runs=5, + logger=None, device=torch.device("cpu"), ): super().__init__() @@ -1099,6 +1100,13 @@ class RPL(Task): self.train_input = self.tensorize(train_sequences) self.test_input = self.tensorize(test_sequences) + if logger is not None: + for x in self.train_input[:25]: + end = (x != self.t_nul).nonzero().max().item() + 1 + seq = [self.id2token[i.item()] for i in x[:end]] + s = " ".join(seq) + logger(f"example_seq {s}") + self.nb_codes = max(self.train_input.max(), self.test_input.max()) + 1 def batches(self, split="train", nb_to_use=-1, desc=None): @@ -1112,7 +1120,7 @@ class RPL(Task): input.split(self.batch_size), dynamic_ncols=True, desc=desc ): last = (batch != self.t_nul).max(0).values.nonzero().max() + 3 - batch = batch[:, :last] + batch = batch[:, :last].to(self.device) yield batch def vocabulary_size(self): @@ -1121,6 +1129,7 @@ class RPL(Task): def produce_results( self, n_epoch, model, result_dir, logger, deterministic_synthesis ): + # -------------------------------------------------------------------- def compute_nb_errors(input, nb_to_log=0): result = input.clone() s = (result == self.t_prog).long() @@ -1147,21 +1156,24 @@ class RPL(Task): _, _, gt_prog, _ = rpl.compute_nb_errors(gt_seq) gt_prog = " ".join([str(x) for x in gt_prog]) prog = " ".join([str(x) for x in prog]) - logger(f"PROG [{gt_prog}] PREDICTED [{prog}]") + comment = "*" if nb_errors == 0 else "-" + logger(f"{comment} PROG [{gt_prog}] PREDICTED [{prog}]") for start_stack, target_stack, result_stack, correct in stacks: - comment = " CORRECT" if correct else "" + comment = "*" if correct else "-" start_stack = " ".join([str(x) for x in start_stack]) target_stack = " ".join([str(x) for x in target_stack]) result_stack = " ".join([str(x) for x in result_stack]) logger( - f" [{start_stack}] -> [{target_stack}] PREDICTED [{result_stack}]{comment}" + f" {comment} [{start_stack}] -> [{target_stack}] PREDICTED [{result_stack}]" ) nb_to_log -= 1 return sum_nb_total, sum_nb_errors + # -------------------------------------------------------------------- + test_nb_total, test_nb_errors = compute_nb_errors( - self.test_input[:1000], nb_to_log=10 + self.test_input[:1000].to(self.device), nb_to_log=10 ) logger(