From 49031a9ff95d0b341ffee2ebccb150181024086e Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Nicolas=20B=C3=A9reux?= Date: Mon, 13 Apr 2026 17:06:47 +0200 Subject: [PATCH 01/16] Stefano's Code --- rbms/optim.py | 138 ++++++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 138 insertions(+) diff --git a/rbms/optim.py b/rbms/optim.py index e2d708e..b12c7d0 100644 --- a/rbms/optim.py +++ b/rbms/optim.py @@ -54,6 +54,140 @@ def step(self, closure=None): return super().step(closure) +class SR_CG(Optimizer): + def __init__( + self, + params, + lr=0.001, + cg_steps=10, + reg=1e-4, + update_freq=1, + warm_start=True, + maximize=True, + ): + defaults = dict( + lr=lr, + cg_steps=cg_steps, + reg=reg, + update_freq=update_freq, + warm_start=warm_start, + maximize=maximize, + step=0, + ) + self._model = params + super().__init__(params, defaults) + + # Initialize state memory for warm starts + for group in self.param_groups: + for p in group["params"]: + self.state[p]["last_dt"] = torch.zeros_like(p.data) + + @torch.no_grad() + def _fvp(self, p_list, v_chain, tanh_term, model, reg, params): + # p_dict = {} + # for p_tensor, param_ref in zip(p_list, params): + # if param_ref is model.weight_matrix: + # p_dict["w"] = p_tensor + # elif param_ref is model.vbias: + # p_dict["v"] = p_tensor + # elif param_ref is model.hbias: + # p_dict["h"] = p_tensor + + O_dot_p = ( + ((v_chain @ model.weight_matrix) * tanh_term).sum(dim=1) + + (v_chain @ model.vbias) + + (tanh_term @ model.hbias) + ) + + O_dot_p_c = O_dot_p - O_dot_p.mean() + + B = v_chain.size(0) + + Sx_w = (v_chain.T @ (tanh_term * O_dot_p_c.unsqueeze(1))) / B + Sx_v = (v_chain.T @ O_dot_p_c) / B + Sx_h = (tanh_term.T @ O_dot_p_c) / B + + Sx_list = [] + for p_tensor, param_ref in zip(p_list, params): + if param_ref is model.weight_matrix: + Sx_list.append(Sx_w + reg * p_tensor) + elif param_ref is model.vbias: + Sx_list.append(Sx_v + reg * p_tensor) + elif param_ref is model.hbias: + Sx_list.append(Sx_h + reg * p_tensor) + + return Sx_list + + @torch.no_grad() + def step(self, v_chain, closure=None): + for group in self.param_groups: + group["step"] += 1 + params = group["params"] + lr = group["lr"] + reg = group["reg"] + + g = [p.grad.clone() for p in params] + + # Lazy Preconditioning: Only run CG every 'update_freq' steps + if group["step"] % group["update_freq"] == 0 or group["step"] == 1: + local_field = self._model.hbias + v_chain @ self._model.weight_matrix + tanh_term = torch.tanh(local_field) + + if group["warm_start"] and group["step"] > 1: + # Warm Start: Initialize with previous Natural Gradient + delta_theta = [self.state[p]["last_dt"].clone() for p in params] + S_dt = self._fvp( + delta_theta, v_chain, tanh_term, self._model, reg, params + ) + residual = [g_i - s_i for g_i, s_i in zip(g, S_dt)] + else: + # Cold Start + delta_theta = [torch.zeros_like(p) for p in params] + residual = [g_i.clone() for g_i in g] + + p_vec = [r_i.clone() for r_i in residual] + r_dot_r = sum(torch.sum(r * r) for r in residual) + + for _ in range(group["cg_steps"]): + S_p = self._fvp(p_vec, v_chain, tanh_term, self._model, reg, params) + p_Sp = sum(torch.sum(pv * spv).item() for pv, spv in zip(p_vec, S_p)) + + if p_Sp <= 1e-8: + break + + alpha = r_dot_r / p_Sp + + for dt, pv in zip(delta_theta, p_vec): + dt.add_(pv, alpha=alpha) + + for r, spv in zip(residual, S_p): + r.sub_(spv, alpha=alpha) + + new_r_dot_r = sum(torch.sum(r * r).item() for r in residual) + + if new_r_dot_r < 1e-6: + break + + beta = new_r_dot_r / r_dot_r + + for pv, r in zip(p_vec, residual): + pv.mul_(beta).add_(r) + + r_dot_r = new_r_dot_r + + # Save the Natural Gradient for the next warm start + for p_tensor, dt in zip(params, delta_theta): + self.state[p_tensor]["last_dt"].copy_(dt) + else: + # Fallback: Standard Gradient Update for intermediate steps + delta_theta = g + + # Apply the update + direction = 1 if group["maximize"] else -1 + for p_tensor, dt in zip(params, delta_theta): + p_tensor.add_(dt, alpha=direction * lr) + + def setup_optim(optim: str, args: dict, params: EBM) -> list[Optimizer]: match args["optim"]: case "sgd": @@ -62,6 +196,10 @@ def setup_optim(optim: str, args: dict, params: EBM) -> list[Optimizer]: optim_class = SGD_cossim case "adam": optim_class = Adam + case "sr": + optim_class = lambda p, **kwargs: SR_CG( + p, update_freq=1, warm_start=True, **kwargs + ) case _: print(f"Unrecognized optimizer {args['optim']}, falling back to SGD.") optim_class = SGD From 5ee93b12d4aaedfe0056b7fff21ecaed14e6b8c0 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Nicolas=20B=C3=A9reux?= Date: Mon, 20 Apr 2026 16:58:56 +0200 Subject: [PATCH 02/16] Fix hessian --- rbms/bernoulli_bernoulli/classes.py | 15 ++++++++++++ rbms/classes.py | 3 +++ rbms/ising_ising/classes.py | 8 ++++++ rbms/optim.py | 38 +++++++++++++++++------------ rbms/sampler/cd.py | 3 +++ rbms/sampler/pcd.py | 3 +++ rbms/sampler/rdm.py | 3 +++ rbms/scripts/train_rbm.py | 22 ++++++++--------- 8 files changed, 68 insertions(+), 27 deletions(-) diff --git a/rbms/bernoulli_bernoulli/classes.py b/rbms/bernoulli_bernoulli/classes.py index e9e7656..2ea516e 100644 --- a/rbms/bernoulli_bernoulli/classes.py +++ b/rbms/bernoulli_bernoulli/classes.py @@ -16,6 +16,7 @@ ) from rbms.classes import RBM from rbms.custom_fn import check_keys_dict +from rbms.dataset.utils import get_covariance_matrix class BBRBM(RBM): @@ -121,6 +122,20 @@ def compute_gradient(self, data, chains, centered=True): weight_matrix=self.weight_matrix, centered=centered, ) + # self.vbias.grad = torch.zeros_like(self.vbias) + tmp = torch.mean(data["visible"], 0) @ self.weight_matrix + self.hbias.grad = ( + torch.inverse( + (self.weight_matrix.T @ self.weight_matrix) + - tmp.unsqueeze(1) @ tmp.unsqueeze(0) + ) + @ self.hbias.grad + ) + + self.weight_matrix.grad = ( + get_covariance_matrix(data["visible"], data["weights"]).cuda() + @ torch.pinverse(self.weight_matrix).cuda().T + ) def independent_model(self): return BBRBM( diff --git a/rbms/classes.py b/rbms/classes.py index aaaeee2..bc8ea69 100644 --- a/rbms/classes.py +++ b/rbms/classes.py @@ -371,3 +371,6 @@ def get_metrics_display( @abstractmethod def get_metrics_save(self) -> dict[str, np.ndarray] | None: ... + + @abstractmethod + def get_curr_conf(self) -> dict[str, Tensor]: ... diff --git a/rbms/ising_ising/classes.py b/rbms/ising_ising/classes.py index dce007a..a41b361 100644 --- a/rbms/ising_ising/classes.py +++ b/rbms/ising_ising/classes.py @@ -6,6 +6,7 @@ from rbms.classes import RBM from rbms.custom_fn import check_keys_dict, log2cosh +from rbms.dataset.utils import get_covariance_matrix from rbms.ising_ising.implement import ( _compute_energy, _compute_energy_hiddens, @@ -121,6 +122,13 @@ def compute_gradient(self, data, chains, centered=True): weight_matrix=self.weight_matrix, centered=centered, ) + self.vbias.grad = torch.zeros_like(self.vbias) + self.hbias.grad = torch.zeros_like(self.hbias) + + self.weight_matrix.grad = ( + get_covariance_matrix(data["visible"], data["weights"]).cuda() + @ torch.pinverse(self.weight_matrix).cuda().T + ) def independent_model(self): return IIRBM( diff --git a/rbms/optim.py b/rbms/optim.py index b12c7d0..86371ee 100644 --- a/rbms/optim.py +++ b/rbms/optim.py @@ -5,7 +5,7 @@ from torch import Tensor from torch.optim import SGD, Adam, Optimizer -from rbms.classes import EBM +from rbms.classes import EBM, Sampler class SGD_cossim(SGD): @@ -59,12 +59,16 @@ def __init__( self, params, lr=0.001, + maximize=True, cg_steps=10, reg=1e-4, update_freq=1, warm_start=True, - maximize=True, ): + + self._model = None + self._sampler = None + # super().__init__(params, defaults) defaults = dict( lr=lr, cg_steps=cg_steps, @@ -74,25 +78,20 @@ def __init__( maximize=maximize, step=0, ) - self._model = params super().__init__(params, defaults) - # Initialize state memory for warm starts for group in self.param_groups: for p in group["params"]: self.state[p]["last_dt"] = torch.zeros_like(p.data) + def set_sampler(self, sampler: Sampler): + self._sampler = sampler + + def set_model(self, model: EBM): + self._model = model + @torch.no_grad() def _fvp(self, p_list, v_chain, tanh_term, model, reg, params): - # p_dict = {} - # for p_tensor, param_ref in zip(p_list, params): - # if param_ref is model.weight_matrix: - # p_dict["w"] = p_tensor - # elif param_ref is model.vbias: - # p_dict["v"] = p_tensor - # elif param_ref is model.hbias: - # p_dict["h"] = p_tensor - O_dot_p = ( ((v_chain @ model.weight_matrix) * tanh_term).sum(dim=1) + (v_chain @ model.vbias) @@ -119,7 +118,10 @@ def _fvp(self, p_list, v_chain, tanh_term, model, reg, params): return Sx_list @torch.no_grad() - def step(self, v_chain, closure=None): + def step(self, closure=None): + assert self._sampler is not None + assert self._model is not None + v_chain = self._sampler.get_curr_conf()["visible"] for group in self.param_groups: group["step"] += 1 params = group["params"] @@ -188,7 +190,7 @@ def step(self, v_chain, closure=None): p_tensor.add_(dt, alpha=direction * lr) -def setup_optim(optim: str, args: dict, params: EBM) -> list[Optimizer]: +def setup_optim(optim: str, args: dict, params: EBM, sampler: Sampler) -> list[Optimizer]: match args["optim"]: case "sgd": optim_class = SGD @@ -245,5 +247,9 @@ def setup_optim(optim: str, args: dict, params: EBM) -> list[Optimizer]: ) for opt in optimizer ] - + if args["optim"] == "sr": + for opt in optimizer: + assert isinstance(opt, SR_CG) + opt.set_sampler(sampler) + opt.set_model(params) return optimizer diff --git a/rbms/sampler/cd.py b/rbms/sampler/cd.py index 5945b01..bff6c4b 100644 --- a/rbms/sampler/cd.py +++ b/rbms/sampler/cd.py @@ -18,6 +18,9 @@ def get_conf_grad(self, batch: Tensor) -> dict[str, Tensor]: self.sample(num_steps=None, batch=batch) return self.chains + def get_curr_conf(self): + return self.chains + def sample(self, num_steps: int | None, **kwargs) -> None: batch = kwargs["batch"] self.chains = self.params.init_chains(num_samples=batch.shape[0], start_v=batch) diff --git a/rbms/sampler/pcd.py b/rbms/sampler/pcd.py index 7a52eb4..205e437 100644 --- a/rbms/sampler/pcd.py +++ b/rbms/sampler/pcd.py @@ -27,6 +27,9 @@ def get_conf_grad(self, batch: Tensor): self.sample(num_steps=None) return self.chains + def get_curr_conf(self): + return self.chains + def sample(self, num_steps: int | None, **kwargs): self.chains = self.params.sample_state( chains=self.chains, n_steps=self.num_steps, beta=self.beta diff --git a/rbms/sampler/rdm.py b/rbms/sampler/rdm.py index f8dd845..c0e3436 100644 --- a/rbms/sampler/rdm.py +++ b/rbms/sampler/rdm.py @@ -28,6 +28,9 @@ def get_conf_grad(self, batch: Tensor): self.sample(num_steps=None) return self.chains + def get_curr_conf(self): + return self.chains + @torch.compiler.disable def named_parameters(self): params_dict = self.params.named_parameters() diff --git a/rbms/scripts/train_rbm.py b/rbms/scripts/train_rbm.py index 1c65c36..57a5fee 100644 --- a/rbms/scripts/train_rbm.py +++ b/rbms/scripts/train_rbm.py @@ -2,6 +2,7 @@ import h5py import torch +from torch.optim.optimizer import Optimizer from rbms import get_saved_updates from rbms.dataset import load_dataset @@ -119,17 +120,6 @@ def main(args, map_model=map_model): map_model=map_model, ) - optimizer = setup_optim(args["optim"], args, params) - from rbms.pre_grad import build_pre_grad_update - - pre_grad_update = build_pre_grad_update( - optimizer=optimizer, - lambda_l1=args["L1"], - lambda_l2=args["L2"], - normalize_grad=args["normalize_grad"], - max_grad_norm=args["max_norm_grad"], - ) - match args["training_type"]: case "pcd": sampler = PCD( @@ -151,6 +141,16 @@ def main(args, map_model=map_model): case _: raise ValueError(f"No training type {args['training_type']} supported.") + optimizer: list[Optimizer] = setup_optim(args["optim"], args, params, sampler) + from rbms.pre_grad import build_pre_grad_update + + pre_grad_update = build_pre_grad_update( + optimizer=optimizer, + lambda_l1=args["L1"], + lambda_l2=args["L2"], + normalize_grad=args["normalize_grad"], + max_grad_norm=args["max_norm_grad"], + ) train( train_dataset=train_dataset, test_dataset=test_dataset, From b41d86a55a6ddafb82282b7e3b566bfd06b8337c Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Nicolas=20B=C3=A9reux?= Date: Mon, 18 May 2026 18:17:54 +0200 Subject: [PATCH 03/16] relax pytorch version dependency --- pyproject.toml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/pyproject.toml b/pyproject.toml index 16a6890..872dafc 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -24,7 +24,7 @@ dependencies = [ "h5py>=3.14.0", "numpy>=2.0.0", "matplotlib>=3.8.0", - "torch>=2.10.0", + "torch>=2.8.0", "tqdm>=4.65.0", ] From 6d049c6985538634dab20266f7371113d21ba4ce Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Nicolas=20B=C3=A9reux?= Date: Fri, 29 May 2026 18:14:11 +0200 Subject: [PATCH 04/16] clean import placement --- rbms/scripts/train_rbm.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/rbms/scripts/train_rbm.py b/rbms/scripts/train_rbm.py index 57a5fee..58699e3 100644 --- a/rbms/scripts/train_rbm.py +++ b/rbms/scripts/train_rbm.py @@ -21,6 +21,7 @@ remove_argument, set_args_default, ) +from rbms.pre_grad import build_pre_grad_update from rbms.sampler import CD, PCD, RDM from rbms.training.implement import _init_training, _restore_training from rbms.training.pcd import train @@ -142,7 +143,6 @@ def main(args, map_model=map_model): raise ValueError(f"No training type {args['training_type']} supported.") optimizer: list[Optimizer] = setup_optim(args["optim"], args, params, sampler) - from rbms.pre_grad import build_pre_grad_update pre_grad_update = build_pre_grad_update( optimizer=optimizer, From 03781a9237d761e3a55661cd207565c4df623180 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Nicolas=20B=C3=A9reux?= Date: Fri, 29 May 2026 18:14:58 +0200 Subject: [PATCH 05/16] split data keep variable type --- rbms/scripts/split_data.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/rbms/scripts/split_data.py b/rbms/scripts/split_data.py index 3783702..013e6f2 100644 --- a/rbms/scripts/split_data.py +++ b/rbms/scripts/split_data.py @@ -137,11 +137,13 @@ def split_data_train_test( with h5py.File(output_train_file, "w") as f: f["samples"] = data_train f["labels"] = labels_train + f["variable_type"] = dataset.variable_type print(" Done") print(f"Writing test dataset to '{output_test_file}'...") with h5py.File(output_test_file, "w") as f: f["samples"] = data_test f["labels"] = labels_test + f["variable_type"] = dataset.variable_type print(" Done") case "fasta": From a912026f28b6027a84dbcb328b08b214f64bc455 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Nicolas=20B=C3=A9reux?= Date: Fri, 29 May 2026 18:15:30 +0200 Subject: [PATCH 06/16] plot image add cmap and figsize option --- rbms/plot.py | 12 +++++++++--- 1 file changed, 9 insertions(+), 3 deletions(-) diff --git a/rbms/plot.py b/rbms/plot.py index 2293154..6adab4b 100644 --- a/rbms/plot.py +++ b/rbms/plot.py @@ -110,7 +110,13 @@ def plot_PCA(data1, data2, labels, dir1=0, dir2=1, log=False): def plot_image( - sample, shape=(28, 28), grid_size=(10, 10), show_grid=False, randomize=True + sample, + shape=(28, 28), + grid_size=(10, 10), + show_grid=False, + randomize=True, + cmap="gray", + figsize=(5, 5), ): """Args: sample @@ -142,8 +148,8 @@ def plot_image( ] = sample[id_s].reshape(shape) # Directly reshape to `shape` # Plot the display image - fig, ax = plt.subplots(1, 1) - ax.imshow(display, cmap="gray") + fig, ax = plt.subplots(1, 1, figsize=figsize) + ax.imshow(display, cmap=cmap) ax.axis("off") # Hide axes if show_grid: From 2f28ec5f45e625a3b107911460769d03dd5425ba Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Nicolas=20B=C3=A9reux?= Date: Fri, 29 May 2026 18:19:34 +0200 Subject: [PATCH 07/16] verbose args dataset --- rbms/dataset/__init__.py | 4 +++- rbms/dataset/dataset_class.py | 5 +++-- 2 files changed, 6 insertions(+), 3 deletions(-) diff --git a/rbms/dataset/__init__.py b/rbms/dataset/__init__.py index bf96690..c98b53c 100644 --- a/rbms/dataset/__init__.py +++ b/rbms/dataset/__init__.py @@ -18,6 +18,7 @@ def load_dataset( remove_duplicates: bool = False, device: torch.device | str = "cpu", dtype: torch.dtype = torch.float32, + verbose: bool = True, ) -> tuple[RBMDataset, RBMDataset | None]: return_datasets = [] for dset_name in [dataset_name, test_dataset_name]: @@ -28,7 +29,8 @@ def load_dataset( if dset_name is not None: dset_name = Path(dset_name) - print(f"Reading dataset from {str(dset_name)}...") + if verbose: + print(f"Reading dataset from {str(dset_name)}...") match dset_name.suffix: case ".h5": data, labels, variable_type, weights = load_HDF5( diff --git a/rbms/dataset/dataset_class.py b/rbms/dataset/dataset_class.py index d73aa69..b9cda66 100644 --- a/rbms/dataset/dataset_class.py +++ b/rbms/dataset/dataset_class.py @@ -1,4 +1,5 @@ from __future__ import annotations + import gzip import textwrap from typing import Union @@ -128,9 +129,9 @@ def get_gzip_entropy(self, mean_size: int = 50, num_samples: int = 100): ) return np.mean(en) - def match_model_variable_type(self, visible_type: str): + def match_model_variable_type(self, visible_type: str, verbose: bool = True): self.data = convert_data[self.variable_type][visible_type](self.data) - if self.variable_type != visible_type: + if self.variable_type != visible_type and verbose: print(f"Converting from '{self.variable_type}' to '{visible_type}'") print(self.data) self.variable_type = visible_type From 8d3b52b3244feb6ff647edcd88ea7c840b704e39 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Nicolas=20B=C3=A9reux?= Date: Fri, 29 May 2026 18:28:55 +0200 Subject: [PATCH 08/16] SR_CG -> NGD optimizer --- rbms/optim.py | 355 +++++++++++++++++++++++++++++++++----------------- 1 file changed, 234 insertions(+), 121 deletions(-) diff --git a/rbms/optim.py b/rbms/optim.py index 86371ee..b1baac7 100644 --- a/rbms/optim.py +++ b/rbms/optim.py @@ -54,35 +54,105 @@ def step(self, closure=None): return super().step(closure) -class SR_CG(Optimizer): +def setup_optim(optim: str, args: dict, params: EBM, sampler: Sampler) -> list[Optimizer]: + match args["optim"]: + case "sgd": + optim_class = SGD + case "cossim": + optim_class = SGD_cossim + case "adam": + optim_class = Adam + case "ngd": + optim_class = lambda p, **kwargs: NGD( + p, update_freq=1, warm_start=True, update_biases=True, **kwargs + ) + case _: + print(f"Unrecognized optimizer {args['optim']}, falling back to SGD.") + optim_class = SGD + learning_rate = args["learning_rate"] + max_lr = args["max_lr"] + if args["scale_lr"]: + learning_rate /= np.sqrt(params.effective_number_variables) + max_lr /= np.sqrt(params.effective_number_variables) + + if args["mult_optim"]: + if not isinstance(learning_rate, Tensor): + learning_rate = torch.tensor([learning_rate] * len(params.parameters())) + optimizer = [ + optim_class( + [p], + lr=learning_rate[i], + maximize=True, + ) + for i, p in enumerate(params.parameters()) + ] + else: + if not isinstance(learning_rate, Tensor): + learning_rate = torch.tensor([learning_rate]) + optimizer = [ + optim_class( + params.parameters(), + lr=learning_rate[0], + maximize=True, + ) + ] + for opt in optimizer: + if isinstance(opt, SGD_cossim): + opt.max_lr = max_lr + + if args["optim"] == "nag": + optimizer = [ + SGD( + opt.param_groups[0]["params"], + lr=opt.param_groups[0]["lr"], + maximize=True, + momentum=0.9, + nesterov=True, + ) + for opt in optimizer + ] + if args["optim"] == "ngd": + for opt in optimizer: + assert isinstance(opt, NGD) + opt.set_sampler(sampler) + opt.set_model(params) + return optimizer + + +class NGD(Optimizer): + # Added 'update_biases=True' flag to the initialization def __init__( self, params, lr=0.001, - maximize=True, - cg_steps=10, - reg=1e-4, + cg_steps=20, + init_reg=1, update_freq=1, - warm_start=True, + warm_start=False, + maximize=True, + update_biases=True, + max_lr=0.01, ): - self._model = None self._sampler = None - # super().__init__(params, defaults) defaults = dict( lr=lr, cg_steps=cg_steps, - reg=reg, + reg=init_reg, update_freq=update_freq, warm_start=warm_start, maximize=maximize, + update_biases=update_biases, step=0, + max_lr=max_lr, ) super().__init__(params, defaults) - # Initialize state memory for warm starts + for group in self.param_groups: for p in group["params"]: self.state[p]["last_dt"] = torch.zeros_like(p.data) + self.max_lr = max_lr + self.rescale_eigenvalue_lr = True def set_sampler(self, sampler: Sampler): self._sampler = sampler @@ -91,73 +161,136 @@ def set_model(self, model: EBM): self._model = model @torch.no_grad() - def _fvp(self, p_list, v_chain, tanh_term, model, reg, params): - O_dot_p = ( - ((v_chain @ model.weight_matrix) * tanh_term).sum(dim=1) - + (v_chain @ model.vbias) - + (tanh_term @ model.hbias) - ) + def _get_adaptive_reg( + self, v_chain_eff, tanh_term, model, base_reg + ): # <-- Use v_chain_eff + B = v_chain_eff.size(0) - O_dot_p_c = O_dot_p - O_dot_p.mean() + # Flatten the 3D one-hot tensor to 2D + v_flat = v_chain_eff.view(B, -1) + + v_sq_sum = (v_flat**2).sum(dim=1) + t_sq_sum = (tanh_term**2).sum(dim=1) + + mean_norm_s_k_sq = (v_sq_sum * t_sq_sum).mean() + mean_s_w = (v_flat.T @ tanh_term) / B # Now safe! + norm_mean_s_sq = (mean_s_w**2).sum() + + trace_F = mean_norm_s_k_sq - norm_mean_s_sq + D = model.weight_matrix.numel() + + adaptive_reg = base_reg * trace_F.item() / (2 * D) + return adaptive_reg # Ensure strict positivity + @torch.no_grad() + def _fvp(self, p_list, v_chain, tanh_term, reg, update_biases): B = v_chain.size(0) - Sx_w = (v_chain.T @ (tanh_term * O_dot_p_c.unsqueeze(1))) / B - Sx_v = (v_chain.T @ O_dot_p_c) / B - Sx_h = (tanh_term.T @ O_dot_p_c) / B + # 1. Flatten v_chain to 2D: (B, N_v * N_s) + v_chain_flat = v_chain.view(B, -1) + + if update_biases: + p_w, p_v, p_h = p_list + + # Flatten weights and biases to match + p_w_flat = p_w.view(-1, p_w.size(-1)) # (N_v * N_s, N_h) + p_v_flat = p_v.view(-1) # (N_v * N_s,) + + O_dot_p = ( + ((v_chain_flat @ p_w_flat) * tanh_term).sum(dim=1) + + (v_chain_flat @ p_v_flat) + + (tanh_term @ p_h) + ) + else: + p_w = p_list[0] + p_w_flat = p_w.view(-1, p_w.size(-1)) + O_dot_p = ((v_chain_flat @ p_w_flat) * tanh_term).sum(dim=1) + + O_dot_p_c = O_dot_p - O_dot_p.mean() - Sx_list = [] - for p_tensor, param_ref in zip(p_list, params): - if param_ref is model.weight_matrix: - Sx_list.append(Sx_w + reg * p_tensor) - elif param_ref is model.vbias: - Sx_list.append(Sx_v + reg * p_tensor) - elif param_ref is model.hbias: - Sx_list.append(Sx_h + reg * p_tensor) + # 2. Compute flat gradients + Sx_w_flat = (v_chain_flat.T @ (tanh_term * O_dot_p_c.unsqueeze(1))) / B - return Sx_list + # 3. Reshape back to original parameter shape using .view_as() + Sx_w = Sx_w_flat.view_as(p_w) + + if update_biases: + Sx_v_flat = (v_chain_flat.T @ O_dot_p_c) / B + Sx_v = Sx_v_flat.view_as(p_v) + Sx_h = (tanh_term.T @ O_dot_p_c) / B + return [Sx_w + reg * p_w, Sx_v + reg * p_v, Sx_h + reg * p_h] + else: + return [Sx_w + reg * p_w] @torch.no_grad() - def step(self, closure=None): - assert self._sampler is not None + def step(self, scale=1, closure=None): assert self._model is not None + assert self._sampler is not None v_chain = self._sampler.get_curr_conf()["visible"] + # v_chain = self._model.sample_state(self._sampler.get_curr_conf(), 10)["visible"] for group in self.param_groups: group["step"] += 1 params = group["params"] lr = group["lr"] - reg = group["reg"] + update_biases = group["update_biases"] g = [p.grad.clone() for p in params] - # Lazy Preconditioning: Only run CG every 'update_freq' steps if group["step"] % group["update_freq"] == 0 or group["step"] == 1: - local_field = self._model.hbias + v_chain @ self._model.weight_matrix - tanh_term = torch.tanh(local_field) - + grad_v, grad_h, _ = self._model.compute_energy_visible_gradient(v_chain) + + v_chain_eff = -grad_v + tanh_term = -grad_h + self.reg = scale * self._get_adaptive_reg( + v_chain, tanh_term, self._model, group["reg"] + ) + + # FLAG LOGIC: Filter parameters fed to the Conjugate Gradient solver + if update_biases: + active_params = params + active_grads = g + else: + active_params = [p for p in params if p is self._model.weight_matrix] + active_grads = [p.grad.clone() for p in active_params] + assert len(active_params) > 0 if group["warm_start"] and group["step"] > 1: - # Warm Start: Initialize with previous Natural Gradient - delta_theta = [self.state[p]["last_dt"].clone() for p in params] + delta_theta = [ + self.state[p]["last_dt"].clone() for p in active_params + ] S_dt = self._fvp( - delta_theta, v_chain, tanh_term, self._model, reg, params + delta_theta, v_chain_eff, tanh_term, self.reg, update_biases ) - residual = [g_i - s_i for g_i, s_i in zip(g, S_dt)] + residual = [g_i - s_i for g_i, s_i in zip(active_grads, S_dt)] else: - # Cold Start - delta_theta = [torch.zeros_like(p) for p in params] - residual = [g_i.clone() for g_i in g] + delta_theta = [torch.zeros_like(p) for p in active_params] + residual = [g_i.clone() for g_i in active_grads] p_vec = [r_i.clone() for r_i in residual] r_dot_r = sum(torch.sum(r * r) for r in residual) + # Fix: Do not add an artificial 1e-8. If it's exactly 0, replace it with machine epsilon. + initial_r_norm = torch.sqrt(r_dot_r).item() + if initial_r_norm == 0: + initial_r_norm = 1e-20 + + current_r_norm = initial_r_norm + self.cg_step = 0 + + # while (current_r_norm / initial_r_norm) > 0.001: + # if self.cg_step >= group["cg_steps"]: + # # print("CG broke due to reaching max iterations.") + # break for _ in range(group["cg_steps"]): - S_p = self._fvp(p_vec, v_chain, tanh_term, self._model, reg, params) - p_Sp = sum(torch.sum(pv * spv).item() for pv, spv in zip(p_vec, S_p)) + S_p = self._fvp( + p_vec, v_chain_eff, tanh_term, self.reg, update_biases + ) + p_Sp = sum(torch.sum(pv * spv) for pv, spv in zip(p_vec, S_p)) - if p_Sp <= 1e-8: + if p_Sp.item() <= 1e-20: + print("CG broke due to non-positive curvature.") break - alpha = r_dot_r / p_Sp + alpha = (r_dot_r / p_Sp).item() for dt, pv in zip(delta_theta, p_vec): dt.add_(pv, alpha=alpha) @@ -165,91 +298,71 @@ def step(self, closure=None): for r, spv in zip(residual, S_p): r.sub_(spv, alpha=alpha) - new_r_dot_r = sum(torch.sum(r * r).item() for r in residual) + new_r_dot_r = sum(torch.sum(r * r) for r in residual) + current_r_norm = torch.sqrt(new_r_dot_r).item() - if new_r_dot_r < 1e-6: + if new_r_dot_r.item() < 1e-20: + print("CG broke due to tiny residual norm.") break - beta = new_r_dot_r / r_dot_r + beta = (new_r_dot_r / r_dot_r).item() for pv, r in zip(p_vec, residual): pv.mul_(beta).add_(r) r_dot_r = new_r_dot_r + self.cg_step += 1 - # Save the Natural Gradient for the next warm start - for p_tensor, dt in zip(params, delta_theta): + for p_tensor, dt in zip(active_params, delta_theta): self.state[p_tensor]["last_dt"].copy_(dt) + else: - # Fallback: Standard Gradient Update for intermediate steps delta_theta = g - - # Apply the update - direction = 1 if group["maximize"] else -1 - for p_tensor, dt in zip(params, delta_theta): - p_tensor.add_(dt, alpha=direction * lr) - - -def setup_optim(optim: str, args: dict, params: EBM, sampler: Sampler) -> list[Optimizer]: - match args["optim"]: - case "sgd": - optim_class = SGD - case "cossim": - optim_class = SGD_cossim - case "adam": - optim_class = Adam - case "sr": - optim_class = lambda p, **kwargs: SR_CG( - p, update_freq=1, warm_start=True, **kwargs + active_params = params + active_grads = g + # --- Added Cosine Similarity --- + dot_product = sum( + torch.sum(dt * g_i) for dt, g_i in zip(delta_theta, active_grads) ) - case _: - print(f"Unrecognized optimizer {args['optim']}, falling back to SGD.") - optim_class = SGD - learning_rate = args["learning_rate"] - max_lr = args["max_lr"] - if args["scale_lr"]: - learning_rate /= np.sqrt(params.effective_number_variables) - max_lr /= np.sqrt(params.effective_number_variables) - - if args["mult_optim"]: - if not isinstance(learning_rate, Tensor): - learning_rate = torch.tensor([learning_rate] * len(params.parameters())) - optimizer = [ - optim_class( - [p], - lr=learning_rate[i], - maximize=True, - ) - for i, p in enumerate(params.parameters()) - ] - else: - if not isinstance(learning_rate, Tensor): - learning_rate = torch.tensor([learning_rate]) - optimizer = [ - optim_class( - params.parameters(), - lr=learning_rate[0], - maximize=True, - ) - ] - for opt in optimizer: - if isinstance(opt, SGD_cossim): - opt.max_lr = max_lr - - if args["optim"] == "nag": - optimizer = [ - SGD( - opt.param_groups[0]["params"], - lr=opt.param_groups[0]["lr"], - maximize=True, - momentum=0.9, - nesterov=True, - ) - for opt in optimizer - ] - if args["optim"] == "sr": - for opt in optimizer: - assert isinstance(opt, SR_CG) - opt.set_sampler(sampler) - opt.set_model(params) - return optimizer + norm_dt = torch.sqrt(sum(torch.sum(dt**2) for dt in delta_theta)) + norm_g = torch.sqrt(sum(torch.sum(g_i**2) for g_i in active_grads)) + + self.cos_sim = (dot_product / (norm_dt * norm_g + 1e-20)).item() + # ------------------------------- + # --- Raise learning rate based on cosine similarity --- + # Increases lr if the similarity is positive (up to 2x if perfectly aligned) + group["lr"] *= 1 + self.cos_sim * 0.001 + + # if self.cos_sim > 1e-06: + # group["lr"] *= 1.005 + # elif self.cos_sim < -1e-06: + # group["lr"] *= 0.995 + group["lr"] = min(group["lr"], 0.05) + + # Rescale eigenvalues + if self.rescale_eigenvalue_lr: + self.rescale_eigenvalue_lr = ( + torch.lobpcg( + self._model.weight_matrix @ self._model.weight_matrix.T, + k=1, + largest=True, + )[0][0] + < 1 + ) + if self.rescale_eigenvalue_lr: + _, eig_vector = torch.lobpcg( + delta_theta[0] @ delta_theta[0].T, k=1, largest=True + ) + contrib = ((delta_theta[0].T @ eig_vector) @ eig_vector.T).T + rest = delta_theta[0] - contrib + delta_theta[0] = rest + contrib / self._model.num_visibles + + # if self._model.weight_matrix.svd().S[0] < 1: + # U, S, V = delta_theta[0].svd() + # S[0] /= self._model.num_visibles + # delta_theta[0] = U @ torch.diag_embed(S) @ V.T + + # ------------------------------------------------------ + direction = 1 if group["maximize"] else -1 + for p_tensor, dt in zip(active_params, delta_theta): + p_tensor.add_(dt, alpha=direction * group["lr"]) From 8b93e8bea31801d33852c89d3fe2085d5b5d9a21 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Nicolas=20B=C3=A9reux?= Date: Fri, 29 May 2026 18:29:34 +0200 Subject: [PATCH 09/16] remove approximation natural gradient --- rbms/bernoulli_bernoulli/classes.py | 16 +--------------- 1 file changed, 1 insertion(+), 15 deletions(-) diff --git a/rbms/bernoulli_bernoulli/classes.py b/rbms/bernoulli_bernoulli/classes.py index 2ea516e..46ec5eb 100644 --- a/rbms/bernoulli_bernoulli/classes.py +++ b/rbms/bernoulli_bernoulli/classes.py @@ -16,7 +16,6 @@ ) from rbms.classes import RBM from rbms.custom_fn import check_keys_dict -from rbms.dataset.utils import get_covariance_matrix class BBRBM(RBM): @@ -54,6 +53,7 @@ def __init__( self.hbias = hbias.to(device=self.device, dtype=self.dtype) self.name = "BBRBM" self.flags = [] + self.start = None def __add__(self, other): return BBRBM( @@ -122,20 +122,6 @@ def compute_gradient(self, data, chains, centered=True): weight_matrix=self.weight_matrix, centered=centered, ) - # self.vbias.grad = torch.zeros_like(self.vbias) - tmp = torch.mean(data["visible"], 0) @ self.weight_matrix - self.hbias.grad = ( - torch.inverse( - (self.weight_matrix.T @ self.weight_matrix) - - tmp.unsqueeze(1) @ tmp.unsqueeze(0) - ) - @ self.hbias.grad - ) - - self.weight_matrix.grad = ( - get_covariance_matrix(data["visible"], data["weights"]).cuda() - @ torch.pinverse(self.weight_matrix).cuda().T - ) def independent_model(self): return BBRBM( From c57e906b37ae45bc52f43b3afa3bcdf8eba9489c Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Nicolas=20B=C3=A9reux?= Date: Fri, 29 May 2026 18:30:40 +0200 Subject: [PATCH 10/16] remove approximation natural gradient and add grad energy visible computation --- rbms/ising_ising/classes.py | 15 ++++++++------- rbms/ising_ising/implement.py | 13 +++++++++++++ 2 files changed, 21 insertions(+), 7 deletions(-) diff --git a/rbms/ising_ising/classes.py b/rbms/ising_ising/classes.py index a41b361..fce2a07 100644 --- a/rbms/ising_ising/classes.py +++ b/rbms/ising_ising/classes.py @@ -6,11 +6,11 @@ from rbms.classes import RBM from rbms.custom_fn import check_keys_dict, log2cosh -from rbms.dataset.utils import get_covariance_matrix from rbms.ising_ising.implement import ( _compute_energy, _compute_energy_hiddens, _compute_energy_visibles, + _compute_energy_visibles_gradient, _compute_gradient, _init_chains, _init_parameters, @@ -122,12 +122,13 @@ def compute_gradient(self, data, chains, centered=True): weight_matrix=self.weight_matrix, centered=centered, ) - self.vbias.grad = torch.zeros_like(self.vbias) - self.hbias.grad = torch.zeros_like(self.hbias) - self.weight_matrix.grad = ( - get_covariance_matrix(data["visible"], data["weights"]).cuda() - @ torch.pinverse(self.weight_matrix).cuda().T + def compute_energy_visible_gradient(self, v: Tensor) -> tuple[Tensor, Tensor, Tensor]: + return _compute_energy_visibles_gradient( + v=v, + vbias=self.vbias, + hbias=self.hbias, + weight_matrix=self.weight_matrix, ) def independent_model(self): @@ -157,7 +158,7 @@ def init_chains(self, num_samples, weights=None, start_v=None): ) @staticmethod - def init_parameters(num_hiddens, dataset, device, dtype, var_init=0.0001): + def init_parameters(num_hiddens, dataset, device, dtype, var_init=0.001): data = dataset.data # Convert to torch Tensor if necessary if isinstance(data, np.ndarray): diff --git a/rbms/ising_ising/implement.py b/rbms/ising_ising/implement.py index 72bfcc4..c9fb358 100644 --- a/rbms/ising_ising/implement.py +++ b/rbms/ising_ising/implement.py @@ -159,3 +159,16 @@ def _init_parameters( vbias = torch.atanh(frequencies).to(device=device, dtype=dtype) hbias = torch.zeros(num_hiddens, device=device, dtype=dtype) return vbias, hbias, weight_matrix + + +def _compute_energy_visibles_gradient( + v: Tensor, vbias: Tensor, hbias: Tensor, weight_matrix: Tensor +) -> tuple[Tensor, Tensor, Tensor]: + local_field = hbias + (v @ weight_matrix) + tanh_term = torch.tanh(local_field) + + grad_vbias = -v + grad_hbias = -tanh_term + grad_weight_matrix = -v.T @ tanh_term + # grad_weight_matrix = None + return grad_vbias, grad_hbias, grad_weight_matrix From c853f578a2933645e26808d6aeb5751b81629cac Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Nicolas=20B=C3=A9reux?= Date: Wed, 24 Jun 2026 13:16:22 +0200 Subject: [PATCH 11/16] change cossim to be proportional to cossim --- rbms/optim.py | 13 +++++++------ 1 file changed, 7 insertions(+), 6 deletions(-) diff --git a/rbms/optim.py b/rbms/optim.py index ddc4ef4..23591f4 100644 --- a/rbms/optim.py +++ b/rbms/optim.py @@ -45,16 +45,17 @@ def step(self, closure=None): learning_rate = group["lr"] curr_grad = torch.concatenate([p.grad.flatten() for p in params]).flatten() cosine_similarity = curr_grad @ self.prev_grad - if cosine_similarity > 1e-6: - learning_rate *= 1.002 - elif cosine_similarity < -1e-6: - learning_rate *= 0.998 + # if cosine_similarity > 1e-6: + # learning_rate *= 1.002 + # elif cosine_similarity < -1e-6: + # learning_rate *= 0.998 + learning_rate *= 1 + 1 + cosine_similarity * 0.001 group["lr"] = min(self.max_lr, learning_rate) self.prev_grad = curr_grad.clone() return super().step(closure) -def setup_optim(optim: str, args: dict, params: EBM) -> list[Optimizer]: +def setup_optim(optim: str, args: dict, params: EBM, sampler: Sampler) -> list[Optimizer]: match optim: case "sgd": optim_class = SGD @@ -286,7 +287,7 @@ def step(self, scale=1, closure=None): ) p_Sp = sum(torch.sum(pv * spv) for pv, spv in zip(p_vec, S_p)) - if p_Sp.item() <= 1e-20: + if p_Sp.item() <= 1e-9: print("CG broke due to non-positive curvature.") break From 49edfce2eeb0bed1cc583ed7b796412e6719e0ad Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Nicolas=20B=C3=A9reux?= Date: Wed, 24 Jun 2026 13:35:45 +0200 Subject: [PATCH 12/16] pass effective time as arg --- rbms/training/pcd.py | 1 + 1 file changed, 1 insertion(+) diff --git a/rbms/training/pcd.py b/rbms/training/pcd.py index 4eae5b0..545504b 100644 --- a/rbms/training/pcd.py +++ b/rbms/training/pcd.py @@ -104,6 +104,7 @@ def train( num_updates=idx, time=curr_time + elapsed_time, learning_rate=learning_rate, + effective_time=effective_time, flags=flags, ) From 7d6f8febae8a8fd560b80d0f7dbb2cdad1e31ae0 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Nicolas=20B=C3=A9reux?= Date: Wed, 24 Jun 2026 13:40:59 +0200 Subject: [PATCH 13/16] pass effective time as arg --- rbms/optim.py | 2 +- rbms/training/implement.py | 5 +++-- 2 files changed, 4 insertions(+), 3 deletions(-) diff --git a/rbms/optim.py b/rbms/optim.py index 23591f4..44743e3 100644 --- a/rbms/optim.py +++ b/rbms/optim.py @@ -49,7 +49,7 @@ def step(self, closure=None): # learning_rate *= 1.002 # elif cosine_similarity < -1e-6: # learning_rate *= 0.998 - learning_rate *= 1 + 1 + cosine_similarity * 0.001 + learning_rate *= 1 + cosine_similarity.item() * 0.001 group["lr"] = min(self.max_lr, learning_rate) self.prev_grad = curr_grad.clone() return super().step(closure) diff --git a/rbms/training/implement.py b/rbms/training/implement.py index 9e798dd..a720d15 100644 --- a/rbms/training/implement.py +++ b/rbms/training/implement.py @@ -1,13 +1,13 @@ import h5py import numpy as np import torch +from torch import Tensor from rbms.classes import EBM from rbms.dataset.dataset_class import RBMDataset from rbms.io import load_model, save_model from rbms.map_model import map_model from rbms.utils import get_saved_updates -from torch import Tensor def _init_training( @@ -82,7 +82,7 @@ def _init_training( hyperparameters["num_hiddens"] = num_hiddens hyperparameters["num_chains"] = num_chains hyperparameters["filename"] = str(filename) - + effective_time = torch.zeros_like(lr).cpu().numpy() save_model( filename=filename, params=params, @@ -90,6 +90,7 @@ def _init_training( num_updates=1, time=0.0, flags=flags, + effective_time=effective_time, learning_rate=lr, ) From 8319028ccf46a3f52efb004d80f7fef145ad6d73 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Nicolas=20B=C3=A9reux?= Date: Wed, 24 Jun 2026 15:31:36 +0200 Subject: [PATCH 14/16] compute_energy_visibles take chains as input instead of v --- rbms/bernoulli_bernoulli/classes.py | 4 ++-- rbms/bernoulli_gaussian/classes.py | 4 ++-- rbms/bm/potts.py | 3 ++- rbms/classes.py | 2 +- rbms/ising_gaussian/classes.py | 6 +++--- rbms/ising_ising/classes.py | 4 ++-- rbms/partition_function/ais.py | 4 ++-- rbms/partition_function/exact.py | 6 +++--- rbms/potts_bernoulli/classes.py | 4 ++-- rbms/sampler/pt.py | 4 ++-- rbms/utils.py | 5 ++++- 11 files changed, 25 insertions(+), 21 deletions(-) diff --git a/rbms/bernoulli_bernoulli/classes.py b/rbms/bernoulli_bernoulli/classes.py index e9e7656..40b5138 100644 --- a/rbms/bernoulli_bernoulli/classes.py +++ b/rbms/bernoulli_bernoulli/classes.py @@ -100,9 +100,9 @@ def compute_energy_hiddens(self, h: Tensor) -> Tensor: weight_matrix=self.weight_matrix, ) - def compute_energy_visibles(self, v: Tensor) -> Tensor: + def compute_energy_visibles(self, chains: dict[str, Tensor]) -> Tensor: return _compute_energy_visibles( - v=v, + v=chains["visible"], vbias=self.vbias, hbias=self.hbias, weight_matrix=self.weight_matrix, diff --git a/rbms/bernoulli_gaussian/classes.py b/rbms/bernoulli_gaussian/classes.py index e634f58..8a85816 100644 --- a/rbms/bernoulli_gaussian/classes.py +++ b/rbms/bernoulli_gaussian/classes.py @@ -112,9 +112,9 @@ def compute_energy_hiddens(self, h: Tensor) -> Tensor: h=h, vbias=self.vbias, hbias=self.hbias, weight_matrix=self.weight_matrix ) - def compute_energy_visibles(self, v: Tensor) -> Tensor: + def compute_energy_visibles(self, chains: dict[str, Tensor]) -> Tensor: return _compute_energy_visibles( - v=v, + v=chains["visible"], vbias=self.vbias, hbias=self.hbias, weight_matrix=self.weight_matrix, diff --git a/rbms/bm/potts.py b/rbms/bm/potts.py index ff7e5d1..9ddfc28 100644 --- a/rbms/bm/potts.py +++ b/rbms/bm/potts.py @@ -74,7 +74,7 @@ def sample_visibles( chains["visible"] = one_hot_gen.argmax(-1).to(self.dtype) return chains - def compute_energy_visibles(self, v: Tensor) -> Tensor: + def compute_energy_visibles(self, chains: dict[str, Tensor]) -> Tensor: """Returns the marginalized energy of the model computed on the visible configurations Args: @@ -83,6 +83,7 @@ def compute_energy_visibles(self, v: Tensor) -> Tensor: Returns: Tensor: The computed energy. """ + v = chains["visible"] L, q = self.bias.shape batch_size = v.shape[0] x_flat = ( diff --git a/rbms/classes.py b/rbms/classes.py index aaaeee2..45aaf7c 100644 --- a/rbms/classes.py +++ b/rbms/classes.py @@ -56,7 +56,7 @@ def sample_visibles( ... @abstractmethod - def compute_energy_visibles(self, v: Tensor) -> Tensor: + def compute_energy_visibles(self, chains: dict[str, Tensor]) -> Tensor: """Returns the marginalized energy of the model computed on the visible configurations Args: diff --git a/rbms/ising_gaussian/classes.py b/rbms/ising_gaussian/classes.py index 93c3a2c..ad71790 100644 --- a/rbms/ising_gaussian/classes.py +++ b/rbms/ising_gaussian/classes.py @@ -7,7 +7,7 @@ from torch import Tensor from rbms.classes import RBM -from rbms.custom_fn import check_keys_dict, log2cosh +from rbms.custom_fn import log2cosh from rbms.ising_gaussian.implement import ( _compute_energy, _compute_energy_hiddens, @@ -105,9 +105,9 @@ def compute_energy_hiddens(self, h: Tensor) -> Tensor: h=h, vbias=self.vbias, hbias=self.hbias, weight_matrix=self.weight_matrix ) - def compute_energy_visibles(self, v: Tensor) -> Tensor: + def compute_energy_visibles(self, chains: dict[str, Tensor]) -> Tensor: return _compute_energy_visibles( - v=v, + v=chains["visible"], vbias=self.vbias, hbias=self.hbias, weight_matrix=self.weight_matrix, diff --git a/rbms/ising_ising/classes.py b/rbms/ising_ising/classes.py index dce007a..fc17891 100644 --- a/rbms/ising_ising/classes.py +++ b/rbms/ising_ising/classes.py @@ -100,9 +100,9 @@ def compute_energy_hiddens(self, h: Tensor) -> Tensor: weight_matrix=self.weight_matrix, ) - def compute_energy_visibles(self, v: Tensor) -> Tensor: + def compute_energy_visibles(self, chains: dict[str, Tensor]) -> Tensor: return _compute_energy_visibles( - v=v, + v=chains["visible"], vbias=self.vbias, hbias=self.hbias, weight_matrix=self.weight_matrix, diff --git a/rbms/partition_function/ais.py b/rbms/partition_function/ais.py index d75df8e..2e68533 100644 --- a/rbms/partition_function/ais.py +++ b/rbms/partition_function/ais.py @@ -26,8 +26,8 @@ def update_weights_ais( Tuple[Tensor, dict[str, Tensor]]: A tuple containing the updated log weights and the updated chains. """ chains = prev_params.sample_state(n_steps=n_steps, chains=chains) - energy_prev = prev_params.compute_energy_visibles(v=chains["visible"]) - energy_curr = curr_params.compute_energy_visibles(v=chains["visible"]) + energy_prev = prev_params.compute_energy_visibles(chains) + energy_curr = curr_params.compute_energy_visibles(chains) log_weights += -energy_curr + energy_prev return log_weights, chains diff --git a/rbms/partition_function/exact.py b/rbms/partition_function/exact.py index 3329276..e5fe1c5 100644 --- a/rbms/partition_function/exact.py +++ b/rbms/partition_function/exact.py @@ -1,7 +1,7 @@ import torch from torch import Tensor -from rbms.classes import RBM, EBM +from rbms.classes import EBM, RBM def compute_partition_function_rbm(params: RBM, all_config: Tensor) -> float: @@ -19,7 +19,7 @@ def compute_partition_function_rbm(params: RBM, all_config: Tensor) -> float: if n_dim_config == n_hidden: energy = params.compute_energy_hiddens(h=all_config) elif n_dim_config == n_visible: - energy = params.compute_energy_visibles(v=all_config) + energy = params.compute_energy_visibles({"visible": all_config}) else: raise ValueError( f"The number of dimension for the configurations '{n_dim_config}' does not match the number of visible '{n_visible}' or the number of hidden '{n_hidden}'" @@ -33,7 +33,7 @@ def compute_partition_function(params: EBM, all_config: Tensor) -> float: n_visible = params.num_visibles n_dim_config = all_config.shape[1] if n_dim_config == n_visible: - energy = params.compute_energy_visibles(v=all_config) + energy = params.compute_energy_visibles({"visible": all_config}) else: raise ValueError( f"The number of dimension for the configurations '{n_dim_config}' does not match the number of visible '{n_visible}'." diff --git a/rbms/potts_bernoulli/classes.py b/rbms/potts_bernoulli/classes.py index dbdbf73..ed9ff3b 100644 --- a/rbms/potts_bernoulli/classes.py +++ b/rbms/potts_bernoulli/classes.py @@ -104,9 +104,9 @@ def compute_energy_hiddens(self, h): weight_matrix=self.weight_matrix, ) - def compute_energy_visibles(self, v): + def compute_energy_visibles(self, chains): return _compute_energy_visibles( - v=v, + v=chains["visible"], vbias=self.vbias, hbias=self.hbias, weight_matrix=self.weight_matrix, diff --git a/rbms/sampler/pt.py b/rbms/sampler/pt.py index 686f057..81a4352 100644 --- a/rbms/sampler/pt.py +++ b/rbms/sampler/pt.py @@ -30,8 +30,8 @@ def swap_configurations( n_chains, L = chains[0]["visible"].shape acc_rate = torch.zeros(inverse_temperatures.shape[0] - 1) for idx in range(inverse_temperatures.shape[0] - 1): - energy_0 = params.compute_energy_visibles(v=chains[idx]["visible"]) - energy_1 = params.compute_energy_visibles(v=chains[idx + 1]["visible"]) + energy_0 = params.compute_energy_visibles(chains[idx]) + energy_1 = params.compute_energy_visibles(chains[idx + 1]) delta_energy = ( -energy_1 * inverse_temperatures[idx] diff --git a/rbms/utils.py b/rbms/utils.py index d9ba483..587f955 100644 --- a/rbms/utils.py +++ b/rbms/utils.py @@ -230,7 +230,10 @@ def compute_log_likelihood( float: Log Likelihood. """ w_normalized = w_data / w_data.sum() - return -(params.compute_energy_visibles(v=v_data) @ w_normalized).item() - log_z + return ( + -(params.compute_energy_visibles({"visible": v_data}) @ w_normalized).item() + - log_z + ) @torch.jit.script From f82939e735d205b49a9a64fbc77fbf170ca473ea Mon Sep 17 00:00:00 2001 From: AurelienDecelle Date: Thu, 25 Jun 2026 14:42:41 +0200 Subject: [PATCH 15/16] adding 1/N to mh in IGRBM and changing const sign for energy_visibles --- rbms/ising_gaussian/implement.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/rbms/ising_gaussian/implement.py b/rbms/ising_gaussian/implement.py index 2f208aa..8df76b5 100644 --- a/rbms/ising_gaussian/implement.py +++ b/rbms/ising_gaussian/implement.py @@ -9,8 +9,8 @@ def _sample_hiddens( v: Tensor, weight_matrix: Tensor, hbias: Tensor, beta: float = 1.0 ) -> Tuple[Tensor, Tensor]: - mh = hbias + (v @ weight_matrix) - h = ( + mh = beta*(hbias + (v @ weight_matrix)) / weight_matrix.shape[0] + h = ( torch.randn_like(mh) / torch.sqrt(torch.ones_like(mh) * weight_matrix.shape[0]) + mh ) @@ -49,7 +49,7 @@ def _compute_energy_visibles( field = v @ vbias t = hbias + (v @ weight_matrix) quad_term = 0.5 * (t * t).sum(1) / float(weight_matrix.shape[0]) - return -field - quad_term + const + return -field - quad_term - const def _compute_energy_hiddens( From 47cfb3f7c40133f8ee4fed1c9613a5f53892b6fd Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Nicolas=20B=C3=A9reux?= Date: Tue, 30 Jun 2026 14:44:37 +0200 Subject: [PATCH 16/16] revert cossim lr update --- rbms/optim.py | 10 +++++----- 1 file changed, 5 insertions(+), 5 deletions(-) diff --git a/rbms/optim.py b/rbms/optim.py index 44743e3..2b91c13 100644 --- a/rbms/optim.py +++ b/rbms/optim.py @@ -45,11 +45,11 @@ def step(self, closure=None): learning_rate = group["lr"] curr_grad = torch.concatenate([p.grad.flatten() for p in params]).flatten() cosine_similarity = curr_grad @ self.prev_grad - # if cosine_similarity > 1e-6: - # learning_rate *= 1.002 - # elif cosine_similarity < -1e-6: - # learning_rate *= 0.998 - learning_rate *= 1 + cosine_similarity.item() * 0.001 + if cosine_similarity > 1e-6: + learning_rate *= 1.002 + elif cosine_similarity < -1e-6: + learning_rate *= 0.998 + # learning_rate *= 1 + cosine_similarity.item() * 0.001 group["lr"] = min(self.max_lr, learning_rate) self.prev_grad = curr_grad.clone() return super().step(closure)