From 003152518a850d7f1d7d10b56f11a84bccee22d5 Mon Sep 17 00:00:00 2001 From: Jens Ahrensfeld Date: Sat, 30 May 2026 20:43:42 +0200 Subject: [PATCH] [bugfix] - shuffle data each epoch; replace if chains with elif/else in train() Shuffling eliminates systematic gradient bias from fixed mini-batch ordering. elif/else raises ValueError for unrecognised entity types instead of silently calling None and crashing with a cryptic TypeError. Co-Authored-By: Claude Sonnet 4.6 --- src/rbm/train.py | 10 ++++++---- 1 file changed, 6 insertions(+), 4 deletions(-) diff --git a/src/rbm/train.py b/src/rbm/train.py index 0297b95..5935a6d 100644 --- a/src/rbm/train.py +++ b/src/rbm/train.py @@ -203,21 +203,23 @@ def train(entity: Entity, batch: Mat, status: Status): progress = 0 keep_running = True status.on_change(entity) - cd_func = None if entity.type == Entity.Type.GB_RBM: cd_func = cd_gaussian_binary - if entity.type == Entity.Type.BB_RBM: + elif entity.type == Entity.Type.BB_RBM: cd_func = cd_binary_binary - if entity.type == Entity.Type.GG_RBM: + elif entity.type == Entity.Type.GG_RBM: cd_func = cd_gaussian_gaussian - if entity.type == Entity.Type.BG_RBM: + elif entity.type == Entity.Type.BG_RBM: cd_func = cd_binary_gaussian + else: + raise ValueError(f"Unknown entity type: {entity.type}") entity.grad_zero() for epochs in range(params.num_epochs): if not keep_running: break + batch = batch[np.random.permutation(batch.shape[0])] for mini_batch in to_mini_batch(batch, mini_batch_size): # Contrastive divergence learning: calculate gradients dwhv, dbv, dbh = cd_func(entity, mini_batch)