+ c_quizzes = quiz_machine.generate_quizzes(
+ nb_to_create,
+ model_for_generation=model_for_generation,
+ temperature=args.generation_temperature,
+ )
+
+ nb_correct, seq_logproba = quiz_machine.compute_correctness(
+ c_quizzes,
+ models,
+ bidirectional_validation=args.bidirectional_validation,
+ deterministic_validation=args.deterministic_validation,
+ )
+
+ for n, l in zip(nb_correct, seq_logproba):
+ s = " ".join([str(x.item()) for x in l])
+ logp_file.write(f"{n} {s}\n")
+
+ if args.dirty_debug:
+ nb_correct = torch.randint(
+ len(models) + 1, nb_correct.size(), device=c_quizzes.device
+ )
+
+ quizzes_and_nb_correct_records.append((c_quizzes, nb_correct))
+
+ nv = F.one_hot(nb_correct, num_classes=len(models) + 1).sum(0)
+ nv = " ".join([str(x.item()) for x in nv])
+
+ nb_validated = valid_c_quizzes(
+ quizzes_and_nb_correct_records, standard_validity
+ ).size(0)
+
+ log_string(
+ f"keep c_quizzes model {model_for_generation.id} kept {nv} nb_accumulated {nb_validated} / {nb_to_create}"
+ )
+
+ # store the new c_quizzes which have been validated
+
+ new_c_quizzes = valid_c_quizzes(quizzes_and_nb_correct_records, standard_validity)
+
+ quiz_machine.reverse_random_half_in_place(new_c_quizzes)