+ log_string(
+ f"diff_ce {n_epoch} train {acc_train_loss/nb_train_samples} test {acc_test_loss/nb_test_samples}"
+ )
+
+ # -------------------
+ input = task.test_input[:32, : task.height * task.width]
+ targets = task.test_policies[:32]
+ output_gpt = gpt(mygpt.BracketedSequence(input), with_readout=False).x
+ output = model(output_gpt)
+ losses = (-output.log_softmax(-1) * targets + targets.xlogy(targets)).sum(-1)
+ losses = losses * (input == maze.v_empty)
+ losses = losses / losses.max()
+ losses = losses.reshape(-1, args.maze_height, args.maze_width)
+ input = input.reshape(-1, args.maze_height, args.maze_width)
+ maze.save_image(
+ os.path.join(args.result_dir, f"oneshot_{n_epoch:04d}.png"),
+ mazes=input,
+ score_paths=losses,