[faces_sub_image] - add --mini_batch_size CLI arg
Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
@@ -19,13 +19,14 @@ N_IMAGES = 50
|
|||||||
|
|
||||||
|
|
||||||
class TestModel(Model):
|
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)
|
super().__init__(name, work_dir)
|
||||||
self.unit1 = Entity(
|
self.unit1 = Entity(
|
||||||
(N_VIS, N_HID),
|
(N_VIS, N_HID),
|
||||||
EntityParams(do_gaussian_visible=True, do_gaussian_hidden=False),
|
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,
|
TrainingParams(learning_rate=0.001, momentum=0.9, num_epochs=num_epochs,
|
||||||
l1_lambda=l1_lambda)
|
mini_batch_size=mini_batch_size, l1_lambda=l1_lambda)
|
||||||
)
|
)
|
||||||
|
|
||||||
def forward(self, x: Mat) -> Mat:
|
def forward(self, x: Mat) -> Mat:
|
||||||
@@ -178,8 +179,10 @@ if __name__ == '__main__':
|
|||||||
help='Use single-channel grayscale patches (default: false)')
|
help='Use single-channel grayscale patches (default: false)')
|
||||||
ap.add_argument('--l1_lambda', type=float, default=0.0,
|
ap.add_argument('--l1_lambda', type=float, default=0.0,
|
||||||
help='L1 regularisation strength (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)')
|
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()
|
args = ap.parse_args()
|
||||||
|
|
||||||
if args.grayscale:
|
if args.grayscale:
|
||||||
@@ -190,7 +193,8 @@ if __name__ == '__main__':
|
|||||||
prj_name = 'faces_sub_image_gray' if args.grayscale else 'faces_sub_image'
|
prj_name = 'faces_sub_image_gray' if args.grayscale else 'faces_sub_image'
|
||||||
work_dir = 'results'
|
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)
|
model.init(0.001)
|
||||||
|
|
||||||
if args.load_model:
|
if args.load_model:
|
||||||
|
|||||||
Reference in New Issue
Block a user