import os import numpy import cv2 import matplotlib.pyplot as plt from rbm.model import Model from rbm.entity import Entity, EntityParams, TrainingParams from rbm.image import SubImage, normalize from rbm.matrix import Mat, np, convert DATA_DIR = '/media/jens/cifs/bilder/MachineVision/Caltech_WebFaces/' PATCH = 32 STRIDE = 16 # 50% overlap; set equal to PATCH for non-overlapping GRAYSCALE = False # reassigned in __main__ when --grayscale is set N_CH = 1 if GRAYSCALE else 3 N_VIS = N_CH * PATCH * PATCH # 1024 grayscale / 3072 colour N_HID = 128 N_IMAGES = 50 class TestModel(Model): def __init__(self, name: str, work_dir: str = '.', l1_lambda: float = 0.0): 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=1000, mini_batch_size=1000, l1_lambda=l1_lambda) ) def forward(self, x: Mat) -> Mat: return self.unit1.forward(x) def reconstruct(self, x: Mat) -> Mat: return self.unit1.reconstruct(x) def pad_to_stride(img: numpy.ndarray) -> numpy.ndarray: """Pad so (H - PATCH) and (W - PATCH) are divisible by STRIDE (required by SubImage.check).""" h, w = img.shape[:2] ph = (-h) % STRIDE if h >= PATCH else PATCH - h pw = (-w) % STRIDE if w >= PATCH else PATCH - w return numpy.pad(img, ((0, ph), (0, pw), (0, 0)), mode='constant') def extract_patches_numpy(img_chw: numpy.ndarray) -> numpy.ndarray: """Extract overlapping patches (stride=STRIDE) from (C,H,W) array; returns (N, C*P*P).""" C, H, W = img_chw.shape rows = [img_chw[:, x:x+PATCH, y:y+PATCH].flatten() for x in range(0, H - PATCH + 1, STRIDE) for y in range(0, W - PATCH + 1, STRIDE)] return numpy.stack(rows, axis=0) def _img_to_chw(img: numpy.ndarray) -> numpy.ndarray: """Convert cv2 BGR image to (C, H, W) float64 array, respecting GRAYSCALE flag.""" if GRAYSCALE: gray = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY).astype(numpy.float64) / 255.0 return gray[numpy.newaxis] # (1, H, W) img_f = img[:, :, ::-1].astype(numpy.float64) / 255.0 # BGR→RGB return numpy.ascontiguousarray(numpy.transpose(img_f, (2, 0, 1))) # (3, H, W) def load_patches(data_dir: str, n_images: int = N_IMAGES, start: int = 0): """Load images, extract overlapping patches, return normalised cupy array.""" all_patches = [] files = sorted(f for f in os.listdir(data_dir) if f.endswith('.jpg'))[start:start + n_images] for fname in files: img = cv2.imread(os.path.join(data_dir, fname)) if img is None or img.shape[0] < PATCH or img.shape[1] < PATCH: continue img = pad_to_stride(img) img_chw = _img_to_chw(img) all_patches.append(extract_patches_numpy(img_chw)) patches_np = numpy.concatenate(all_patches, axis=0) std = numpy.std(patches_np, axis=1) patches_np = patches_np[std > 0.01] return normalize(np.array(patches_np)) def load_image_patches(path: str): """Extract normalised patches from a single image via SubImage; return (patches, nx_steps, ny_steps).""" sub = SubImage(PATCH, PATCH, STRIDE, STRIDE) img = cv2.imread(path) img = pad_to_stride(img) h, w = img.shape[:2] nx_steps = (h - PATCH) // STRIDE + 1 ny_steps = (w - PATCH) // STRIDE + 1 img_nchw = np.array(_img_to_chw(img)[numpy.newaxis]) patches_nchw = sub(img_nchw) patches = patches_nchw.reshape(-1, N_VIS) return normalize(patches), nx_steps, ny_steps def assemble_overlapping(patch_list, nx_steps: int, ny_steps: int) -> numpy.ndarray: """Reconstruct image from overlapping patches by averaging contributions per pixel.""" H = (nx_steps - 1) * STRIDE + PATCH W = (ny_steps - 1) * STRIDE + PATCH accum = numpy.zeros((N_CH, H, W), dtype=numpy.float64) count = numpy.zeros((1, H, W), dtype=numpy.float64) for k, p in enumerate(patch_list): x = (k // ny_steps) * STRIDE y = (k % ny_steps) * STRIDE arr = convert(p).reshape(N_CH, PATCH, PATCH) accum[:, x:x+PATCH, y:y+PATCH] += arr count[0, x:x+PATCH, y:y+PATCH] += 1.0 img = (accum / numpy.maximum(count, 1)).transpose(1, 2, 0).squeeze() # (H,W,C) or (H,W) lo, hi = img.min(), img.max() return numpy.clip((img - lo) / (hi - lo + 1e-8), 0, 1) def show_filters(entity: Entity, rows: int = 8, cols: int = 16): W = convert(entity.state.w_hv) fig, axes = plt.subplots(rows, cols, figsize=(cols * 1.2, rows * 1.2)) fig.suptitle(f'Learned GB-RBM filters ({N_HID} hidden units, {PATCH}×{PATCH} RGB patches)', fontsize=10) for j in range(rows * cols): ax = axes[j // cols][j % cols] ax.axis('off') if j >= N_HID: continue filt = W[:, j].reshape(N_CH, PATCH, PATCH).transpose(1, 2, 0).squeeze() lo, hi = filt.min(), filt.max() ax.imshow(numpy.clip((filt - lo) / (hi - lo + 1e-8), 0, 1), cmap='gray' if GRAYSCALE else None) plt.tight_layout() def show_reconstructions(model: TestModel, patches: Mat, n_show: int = 10): fig, axes = plt.subplots(2, n_show, figsize=(n_show * 1.5, 3.5)) axes[0, 0].set_ylabel('Original') axes[1, 0].set_ylabel('Recon') fig.suptitle('Patch reconstructions (normalised display)', fontsize=10) err_total = 0.0 cmap = 'gray' if GRAYSCALE else None def to_img(p): a = convert(p).reshape(N_CH, PATCH, PATCH).transpose(1, 2, 0).squeeze() return numpy.clip((a - a.min()) / (a.max() - a.min() + 1e-8), 0, 1) for i in range(n_show): inp = patches[i] recon = model.reconstruct(model.forward(inp)) err_total += float(np.mean(np.abs(inp - recon))) axes[0, i].imshow(to_img(inp), cmap=cmap); axes[0, i].axis('off') axes[1, i].imshow(to_img(recon), cmap=cmap); axes[1, i].axis('off') print(f'Mean patch reconstruction MAE: {err_total / n_show:.4f}') plt.tight_layout() def show_image_reconstruction(model: TestModel, img_path: str): patches, nx_steps, ny_steps = load_image_patches(img_path) recons = [model.reconstruct(model.forward(patches[i])) for i in range(patches.shape[0])] orig_grid = assemble_overlapping([patches[i] for i in range(len(recons))], nx_steps, ny_steps) recon_grid = assemble_overlapping(recons, nx_steps, ny_steps) fig, axes = plt.subplots(1, 2, figsize=(12, 6)) cmap = 'gray' if GRAYSCALE else None axes[0].imshow(orig_grid, cmap=cmap); axes[0].set_title('Original (normalised)'); axes[0].axis('off') axes[1].imshow(recon_grid, cmap=cmap); axes[1].set_title('Reconstructed (overlap blended)'); axes[1].axis('off') fig.suptitle(os.path.basename(img_path), fontsize=9) plt.tight_layout() if __name__ == '__main__': from argparse import ArgumentParser ap = ArgumentParser() ap.add_argument('--load_model', type=lambda s: s.lower() != 'false', default=True, help='Load saved model weights (default: true)') ap.add_argument('--do_train', type=lambda s: s.lower() != 'false', default=False, help='Train the model (default: false)') ap.add_argument('--grayscale', action='store_true', default=False, 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)') args = ap.parse_args() if args.grayscale: GRAYSCALE = True N_CH = 1 N_VIS = PATCH * PATCH # 1024 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) model.init(0.001) if args.load_model: model.load() if args.do_train: print(f'Loading {N_IMAGES} training images (stride={STRIDE})...') train_patches = load_patches(DATA_DIR, N_IMAGES) print(f'Training on {train_patches.shape[0]} patches ({N_VIS}-dim each)') model.train(train_patches) model.save() # Show learned weight filters show_filters(model.unit1) # Show patch-level reconstructions on held-out images print('Loading test patches...') test_patches = load_patches(DATA_DIR, n_images=10, start=N_IMAGES) if test_patches.shape[0] < 10: test_patches = load_patches(DATA_DIR, n_images=10) show_reconstructions(model, test_patches) # Show full image reconstruction with overlap blending test_img = os.path.join(DATA_DIR, sorted(os.listdir(DATA_DIR))[60]) show_image_reconstruction(model, test_img) plt.show()