Skip to content
AI News HubLIVE
In-site rewrite6 min read

Hierarchical NeRF with JAX3D for Volumetric Rendering, Novel-View Synthesis, and 3D Reconstruction

Summary

In this tutorial, we build an end-to-end hierarchical Neural Radiance Field (NeRF) using JAX, Flax, Optax, and the volume-rendering primitives provided by jax3d. We first construct a synthetic multi-view dataset from an analytic scene containing volumetric geometry and view-dependent radiance, using sample_along_rays and volume_rendering to establish the forward rendering process. We then implement a NeRF […] The post Hierarchical NeRF with JAX3D for Volumetric Rendering, Novel-View Synthesis, and 3D Reconstruction appeared first on MarkTechPost.

SourceMarkTechPostAuthor: Sana Hassan
Hierarchical NeRF with JAX3D for Volumetric Rendering, Novel-View Synthesis, and 3D Reconstruction
Report an error

The correction channel is not available yet. You can copy the article reference below for later.

Correction instructions
Read article

In this tutorial, we build an end-to-end hierarchical Neural Radiance Field (NeRF) using JAX, Flax, Optax, and the volume-rendering primitives provided by jax3d. We first construct a synthetic multi-view dataset from an analytic scene containing volumetric geometry and view-dependent radiance, using sample_along_rays and volume_rendering to establish the forward rendering process. We then implement a NeRF with positional encoding, skip connections, separate coarse and fine networks, and view-direction conditioning, followed by hierarchical importance sampling through sample_piecewise_constant_pdf. We train the model with JAX JIT compilation, Adam optimization, exponential learning-rate decay, and gradient clipping, and finally evaluate novel-view synthesis using PSNR, depth and opacity visualization, sampling diagnostics, 360-degree rendering, and marching-cubes geometry extraction. Copy CodeCopiedUse a different Browser import os, sys, subprocess, importlib.util, functools, dataclasses, time, math def _sh(cmd): subprocess.run(cmd, shell=True, check=False, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL) print("Installing dependencies ...") _sh(f'{sys.executable} -m pip install -q "etils[array-types,epy,etree,enp]" ' f'chex flax optax scikit-image') REPO_DIR = "/content/jax3d" if os.path.isdir("/content") else os.path.abspath("./jax3d") if not os.path.isdir(REPO_DIR): print("Cloning google-research/jax3d ...") _sh(f"git clone -q --depth 1 https://github.com/google-research/jax3d.git {REPO_DIR}") def _load_module_by_path(name, path): """Load a single .py file without triggering the parent package init. from jax3d.math import volume_rendering also works if you run pip install . inside the clone, but that pulls in gin/tfds/etc. """ spec = importlib.util.spec_from_file_location(name, path) mod = importlib.util.module_from_spec(spec) sys.modules[name] = mod spec.loader.exec_module(mod) return mod _VR_PATH = os.path.join(REPO_DIR, "jax3d", "jax3d", "math", "volume_rendering.py") if not os.path.exists(_VR_PATH): _VR_PATH = os.path.join(REPO_DIR, "jax3d", "math", "volume_rendering.py") try: j3vr = _load_module_by_path("j3d_volume_rendering", _VR_PATH) except Exception as e: raise SystemExit( f"Could not load {_VR_PATH}: {e}\n" "Try: pip install -U 'etils[array-types,epy,etree,enp]==1.9.4' and re-run." ) import numpy as np import jax import jax.numpy as jnp import flax.linen as nn import optax from flax.training import train_state import matplotlib.pyplot as plt from PIL import Image print("jax", jax.version, "| device:", jax.devices()[0].device_kind, f"({jax.devices()[0].platform})") print("jax3d volume_rendering API:", [n for n in ("sample_along_rays", "volume_rendering", "sample_piecewise_constant_pdf", "sample_1d") if hasattr(j3vr, n)]) @dataclasses.dataclass class Config: H: int = 64; W: int = 64 n_train_views: int = 24; n_test_views: int = 3 cam_radius: float = 3.2; fov_deg: float = 40.0 near: float = 1.9; far: float = 4.7 gt_samples: int = 256 n_coarse: int = 64; n_fine: int = 64 deg_pos: int = 10; deg_dir: int = 4 width: int = 128; depth: int = 6; skip: int = 3 batch_rays: int = 2048; steps: int = 2500 lr_init: float = 5e-4; lr_final: float = 5e-6 chunk: int = 4096 grid_res: int = 96 cfg = Config() if jax.devices()[0].platform == "cpu": print("\n!! No GPU detected -- switching to a small CPU-friendly config.") print(" (Runtime > Change runtime type > T4 GPU for the full version.)\n") cfg = dataclasses.replace(cfg, H=40, W=40, n_train_views=14, steps=400, gt_samples=128, n_coarse=32, n_fine=32, width=64, depth=4, skip=2, batch_rays=1024, chunk=1600, grid_res=64) def _normalize(v, axis=-1): return v / (np.linalg.norm(v, axis=axis, keepdims=True) + 1e-9) def look_at(eye, target=(0., 0., 0.), up=(0., 0., 1.)): """OpenGL/NeRF convention camera-to-world: +x right, +y up, camera looks at -z.""" eye, target, up = map(lambda a: np.asarray(a, np.float32), (eye, target, up)) fwd = _normalize(target - eye) right = _normalize(np.cross(fwd, up)) trueup = np.cross(right, fwd) c2w = np.eye(4, dtype=np.float32) c2w[:3, :3] = np.stack([right, trueup, -fwd], axis=1) c2w[:3, 3] = eye return c2w def orbit_poses(n, radius, elev_lo=18., elev_hi=58., phase=0.0): """Golden-angle azimuths + monotone elevations => well-spread views on a dome.""" i = np.arange(n, dtype=np.float64) + 0.5 az = 2 * np.pi * ((i * 0.6180339887) + phase) elev = np.arcsin(np.linspace(np.sin(np.deg2rad(elev_lo)), np.sin(np.deg2rad(elev_hi)), n)) eyes = np.stack([radius * np.cos(elev) * np.cos(az), radius * np.cos(elev) * np.sin(az), radius * np.sin(elev)], axis=-1).astype(np.float32) return np.stack([look_at(e) for e in eyes], axis=0) def rays_from_pose(c2w, H, W, focal): """Returns (origins, dirs) of shape [H, W, 3]; dirs are unit-length, so the depths returned by jax3d's sampler are true world-space distances.""" i, j = np.meshgrid(np.arange(W, dtype=np.float32), np.arange(H, dtype=np.float32), indexing="xy") cam_dirs = np.stack([(i - W * .5 + .5) / focal, -(j - H * .5 + .5) / focal, -np.ones_like(i)], axis=-1) dirs = _normalize(cam_dirs @ c2w[:3, :3].T) origins = np.broadcast_to(c2w[:3, 3], dirs.shape) return origins.astype(np.float32).copy(), dirs.astype(np.float32) FOCAL = 0.5 * cfg.W / math.tan(0.5 * math.radians(cfg.fov_deg)) We set up the JAX3D environment, install the required dependencies, and load the volume_rendering module directly from the cloned repository. We configure GPU/CPU-adaptive training parameters and establish the camera model using pinhole intrinsics, look-at poses, and orbit-based camera placement. We then generate normalized world-space rays from each camera pose, providing the geometric foundation for the rendering pipeline. Copy CodeCopiedUse a different Browser LIGHT = jnp.asarray(_normalize(np.array([0.55, 0.75, 0.85], np.float32))) _SPHERES = [ (jnp.array([0.34, 0.02, -0.22]), 0.36, jnp.array([0.90, 0.24, 0.22])), (jnp.array([-0.32, 0.28, 0.05]), 0.26, jnp.array([0.25, 0.78, 0.36])), (jnp.array([-0.05, -0.36, 0.24]), 0.22, jnp.array([0.28, 0.40, 0.95])), ] def _sphere_field(pos, vdir, center, radius, albedo): d = pos - center dist = jnp.linalg.norm(d, axis=-1) n = d / (dist[..., None] + 1e-8) sigma = 80.0 * jax.nn.sigmoid((radius - dist) / 0.015) v = -vdir refl = 2.0 * jnp.sum(n * v, -1, keepdims=True) * n - v spec = 0.65 * jnp.clip(jnp.sum(refl * LIGHT, -1), 0., 1.) 24 lamb = 0.35 + 0.65 * jnp.clip(jnp.sum(n * LIGHT, -1), 0., 1.) rgb = jnp.clip(albedo * lamb[..., None] + spec[..., None], 0., 1.) return sigma, rgb def _floor_field(pos): x, y, z = pos[..., 0], pos[..., 1], pos[..., 2] m = (jax.nn.sigmoid((0.06 - jnp.abs(z + 0.62)) / 0.008) * jax.nn.sigmoid((0.85 - jnp.abs(x)) / 0.01) * jax.nn.sigmoid((0.85 - jnp.abs(y)) / 0.01)) checker = (jnp.floor(x * 3.0) + jnp.floor(y * 3.0)) % 2.0 rgb = jnp.where(checker[..., None] > 0.5, jnp.array([0.86, 0.86, 0.89]), jnp.array([0.22, 0.25, 0.30])) return 80.0 * m, rgb def gt_field(pos, vdir): """pos, vdir: [..., 3] -> (sigma [...], rgb [..., 3]). Density-weighted blend.""" sig_sum = 0.0 col_sum = 0.0 for c, r, a in _SPHERES: s, rgb = _sphere_field(pos, vdir, c, r, a) sig_sum = sig_sum + s col_sum = col_sum + s[..., None] * rgb s, rgb = _floor_field(pos) sig_sum = sig_sum + s col_sum = col_sum + s[..., None] * rgb return sig_sum, col_sum / (sig_sum[..., None] + 1e-8) WHITE_BG = jnp.ones((3,), jnp.float32) @jax.jit def render_ground_truth(origins, dirs): """Fine-grained volumetric render of the analytic scene -> RGB + depth.""" depths, positions = j3vr.sample_along_rays( ray_origins=origins, ray_directions=dirs, near=cfg.near, far=cfg.far, sample_count=cfg.gt_samples, deterministic=True) vdir = jnp.broadcast_to(dirs[..., None, :], positions.shape) sigma, rgb = gt_field(positions, vdir) out = j3vr.volume_rendering( sample_values={"rgb": rgb}, sample_density=sigma, depths=depths, background_values={"rgb": WHITE_BG}) return out.ray_values["rgb"], out.ray_depth, out.ray_alpha def build_dataset(poses): O, D, C = [], [], [] for c2w in poses: o, d = rays_from_pose(c2w, cfg.H, cfg.W, FOCAL) rgb, _, _ = render_ground_truth(jnp.asarray(o), jnp.asarray(d)) O.append(o); D.append(d); C.append(np.asarray(rgb)) return (np.stack(O), np.stack(D), np.stack(C)) print("\nRendering the synthetic multi-view dataset ...") t0 = time.time() train_poses = orbit_poses(cfg.n_train_views, cfg.cam_radius, phase=0.00) test_poses = orbit_poses(cfg.n_test_views, cfg.cam_radius, 26., 50., phase=0.41) tr_o, tr_d, tr_c = build_dataset(train_poses) te_o, te_d, te_c = build_dataset(test_poses) print(f" {cfg.n_train_views} train + {cfg.n_test_views} test views " f"at {cfg.H}x{cfg.W} ({time.time()-t0:.1f}s)") k = min(8, cfg.n_train_views) fig, axes = plt.subplots(1, k, figsize=(2 * k, 2.3)) for a, im, p in zip(axes, tr_c[:k], train_poses[:k]): a.imshow(np.clip(im, 0, 1)); a.axis("off") a.set_title(f"({p[0,3]:+.1f},{p[1,3]:+.1f},{p[2,3]:+.1f})", fontsize=7) fig.suptitle("Training views (ground truth, rendered with jax3d.math.volume_rendering)", fontsize=11); plt.tight_layout(); plt.show() rays_o = jnp.asarray(tr_o.reshape(-1, 3)) rays_d = jnp.asarray(tr_d.reshape(-1, 3)) rays_c = jnp.asarray(tr_c.reshape(-1, 3)) N_RAYS = rays_o.shape[0] print(f" ray pool: {N_RAYS:,} rays") We construct an analytic ground-truth scene containing soft-edged spheres, a patterned floor, and view-dependent specular radiance. We render this scene with JAX3D’s volume-rendering implementation to generate consistent RGB observations, depths, and opacity values across multiple camera views. We organize the resulting images into a flattened ray pool so that we can efficiently sample random rays during NeRF training. Copy CodeCopiedUse a different Browser def posenc(x, deg): """NeRF sinusoidal encoding, with the raw input concatenated.""" if deg == 0: return x scales = 2.0 jnp.arange(deg, dtype=x.dtype) xb = (x[..., None, :] * scales[:, None]).reshape(*x.shape[:-1], -1) return jnp.concatenate([x, jnp.sin(xb), jnp.cos(xb)], axis=-1) class NeRFMLP(nn.Module): width: int; depth: int; skip: int; deg_pos: int; deg_dir: int @nn.compact def call(self, pos, dirs): inp = posenc(pos, self.deg_pos) x = inp for i in range(self.depth): x = nn.relu(nn.Dense(self.width)(x)) if i == self.skip: x = jnp.concatenate([x, inp], axis=-1) sigma = nn.softplus(nn.Dense(1)(x)[..., 0] - 1.0) h = jnp.concatenate([nn.Dense(self.width)(x), posenc(dirs, self.deg_dir)], -1) rgb = nn.sigmoid(nn.Dense(3)(nn.relu(nn.Dense(self.width // 2)(h)))) return sigma, rgb model = NeRFMLP(cfg.width, cfg.depth, cfg.skip, cfg.deg_pos, cfg.deg_dir) We implement the NeRF representation using sinusoidal positional encoding for both spatial coordinates and viewing directions. We use a deep Flax MLP with a skip connection to predict non-negative volumetric density from position while conditioning RGB on the viewing direction. We therefore separate view-independent geometry from view-dependent appearance, allowing the model to represent both scene structure and specular effects. Copy CodeCopiedUse a different Browser def render_rays(params, origins, dirs, rng, deterministic): """Coarse pass -> importance-resample -> fine pass. All sampling and compositing comes from jax3d.math.volume_rendering.""" rng_c, rng_f = jax.random.split(rng) depths_c, pos_c = j3vr.sample_along_rays( ray_origins=origins, ray_directions=dirs, near=cfg.near, far=cfg.far, sample_count=cfg.n_coarse, deterministic=deterministic, rng=rng_c) dirs_c = jnp.broadcast_to(dirs[:, None, :], pos_c.shape) sigma_c, rgb_c = model.apply(params["coarse"], pos_c, dirs_c) out_c = j3vr.volume_rendering( sample_values={"rgb": rgb_c}, sample_density=sigma_c, depths=depths_c, background_values={"rgb": WHITE_BG}) mid = 0.5 * (depths_c[..., 1:] + depths_c[..., :-1]) bin_edges = jnp.concatenate([depths_c[..., :1], mid, depths_c[..., -1:]], -1) t_fine = j3vr.sample_piecewise_constant_pdf( bin_edges=bin_edges, weights=out_c.sample_weights, sample_count=cfg.n_fi [truncated for AI cost control]

Key points and analysis

Article intelligence

EngineersAdvanced

Key points

  • AI generation is temporarily unavailable; this entry was preserved with deterministic fallback metadata.
  • In this tutorial, we build an end-to-end hierarchical Neural Radiance Field (NeRF) using JAX, Flax, Optax, and the volume-rendering primitives provided by jax3d. We first construc…

Highlights and analysis are generated automatically and may contain errors. Check the original source.