From 8e37a868ac7dfc1cb5e924790929c6eebabbeb94 Mon Sep 17 00:00:00 2001 From: =?utf8?q?Fran=C3=A7ois=20Fleuret?= Date: Tue, 25 Jun 2024 23:45:34 +0200 Subject: [PATCH] Update. --- main.py | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/main.py b/main.py index cb28a7d..c5acea7 100755 --- a/main.py +++ b/main.py @@ -394,7 +394,10 @@ def create_c_quizzes( f"keep c_quizzes {to_keep.size(0)}/{new_c_quizzes.size(0)} ({to_keep.size(0)*100/new_c_quizzes.size(0):.02f}%) total {sum([ x.size(0) for x in kept])}/{nb_to_generate}" ) - new_c_quizzes = torch.cat(kept, dim=0)[: nb_for_train + nb_for_test] + new_c_quizzes = torch.cat(kept, dim=0) + new_c_quizzes = new_c_quizzes[ + torch.randperm(new_c_quizzes.size(0))[: nb_for_train + nb_for_test] + ] quizz_machine.store_c_quizzes(new_c_quizzes[:nb_for_train], for_train=True) quizz_machine.store_c_quizzes(new_c_quizzes[nb_for_train:], for_train=False) -- 2.39.5