From 5f5c7c6d77ce85bed92e1d33d033be05ab6302e6 Mon Sep 17 00:00:00 2001 From: Jens Ahrensfeld Date: Fri, 19 Dec 2025 17:29:32 +0100 Subject: [PATCH] Optimizer.train() shall ignore param.do_batch_sample --- src/rbm/train.py | 10 ++++------ 1 file changed, 4 insertions(+), 6 deletions(-) diff --git a/src/rbm/train.py b/src/rbm/train.py index 6177111..a6896b6 100644 --- a/src/rbm/train.py +++ b/src/rbm/train.py @@ -145,17 +145,13 @@ class Optimizer: d_progress = 100.0 / self.params.num_epochs progress = 0 - v_states = batch - if self.params.do_batch_sample: - v_states = sample(batch) - - forward = None for epochs in range(self.params.num_epochs): - forward = self.epoch(v_states) + self.epoch(batch) # check if status update is needed if status.want_report(round(progress)): # Calculate error + forward = self.entity.forward(batch) result = self.entity.reconstruct(forward) err_rms = rms_error_accu(batch - result) if not status.on_change({"progress": {"value": round(progress), "unit": "%"}, @@ -164,6 +160,7 @@ class Optimizer: progress += d_progress + forward = self.entity.forward(batch) err_rms = rms_error_accu(batch - self.entity.reconstruct(forward)) status.on_change({"progress": {"value": round(progress), "unit": "%"}, "err_rms": {"value": err_rms, "unit": ""}}) @@ -174,6 +171,7 @@ class Optimizer: forward = None training_remain = batch.shape[0] batch_row_index = 0 + while training_remain > 0: batch_size = min(self.params.mini_batch_size, training_remain) if batch_size == 0: