From 68b376657c16de5ecac70260dfb4504808f6461e Mon Sep 17 00:00:00 2001 From: Jens Ahrensfeld Date: Sat, 30 May 2026 21:13:25 +0200 Subject: [PATCH] [faces_sub_image] - add --mini_batch_size CLI arg Co-Authored-By: Claude Sonnet 4.6 --- src/tests/test_faces_sub_image.py | 14 +++++++++----- 1 file changed, 9 insertions(+), 5 deletions(-) diff --git a/src/tests/test_faces_sub_image.py b/src/tests/test_faces_sub_image.py index fcac37a..dccb601 100644 --- a/src/tests/test_faces_sub_image.py +++ b/src/tests/test_faces_sub_image.py @@ -19,13 +19,14 @@ N_IMAGES = 50 class TestModel(Model): - def __init__(self, name: str, work_dir: str = '.', l1_lambda: float = 0.0, num_epochs: int = 1000): + def __init__(self, name: str, work_dir: str = '.', l1_lambda: float = 0.0, + num_epochs: int = 1000, mini_batch_size: int = 1000): super().__init__(name, work_dir) self.unit1 = Entity( (N_VIS, N_HID), EntityParams(do_gaussian_visible=True, do_gaussian_hidden=False), - TrainingParams(learning_rate=0.001, momentum=0.9, num_epochs=num_epochs, mini_batch_size=1000, - l1_lambda=l1_lambda) + TrainingParams(learning_rate=0.001, momentum=0.9, num_epochs=num_epochs, + mini_batch_size=mini_batch_size, l1_lambda=l1_lambda) ) def forward(self, x: Mat) -> Mat: @@ -178,8 +179,10 @@ if __name__ == '__main__': help='Use single-channel grayscale patches (default: false)') ap.add_argument('--l1_lambda', type=float, default=0.0, help='L1 regularisation strength (default: 0.0)') - ap.add_argument('--num_epochs', type=int, default=1000, + ap.add_argument('--num_epochs', type=int, default=1000, help='Number of training epochs (default: 1000)') + ap.add_argument('--mini_batch_size', type=int, default=1000, + help='Mini-batch size (default: 1000)') args = ap.parse_args() if args.grayscale: @@ -190,7 +193,8 @@ if __name__ == '__main__': prj_name = 'faces_sub_image_gray' if args.grayscale else 'faces_sub_image' work_dir = 'results' - model = TestModel(prj_name, work_dir, l1_lambda=args.l1_lambda, num_epochs=args.num_epochs) + model = TestModel(prj_name, work_dir, l1_lambda=args.l1_lambda, + num_epochs=args.num_epochs, mini_batch_size=args.mini_batch_size) model.init(0.001) if args.load_model: