{"id":117670,"date":"2026-09-14T01:45:47","date_gmt":"2026-09-14T01:45:47","guid":{"rendered":"https:\/\/youzum.net\/hierarchical-nerf-with-jax3d-for-volumetric-rendering-novel-view-synthesis-and-3d-reconstruction\/"},"modified":"2026-09-14T01:45:47","modified_gmt":"2026-09-14T01:45:47","slug":"hierarchical-nerf-with-jax3d-for-volumetric-rendering-novel-view-synthesis-and-3d-reconstruction","status":"publish","type":"post","link":"https:\/\/youzum.net\/ja\/hierarchical-nerf-with-jax3d-for-volumetric-rendering-novel-view-synthesis-and-3d-reconstruction\/","title":{"rendered":"Hierarchical NeRF with JAX3D for Volumetric Rendering, Novel-View Synthesis, and 3D Reconstruction"},"content":{"rendered":"<p class=\"wp-block-paragraph\">In this<strong><a href=\"https:\/\/github.com\/MARKTECHPOST-AI-MEDIA-INC\/AI-Agents-Projects-Tutorials\/blob\/main\/Computer%20Vision\/jax3d_hierarchical_nerf_tutorial_Marktechpost.ipynb\" target=\"_blank\" rel=\"noreferrer noopener\"> tutorial<\/a><\/strong>, we build an end-to-end hierarchical Neural Radiance Field (NeRF) using <a href=\"https:\/\/github.com\/google-research\/jax3d\"><strong>JAX<\/strong><\/a>, 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.<\/p>\n<div class=\"dm-code-snippet dark dm-normal-version default no-background-mobile\">\n<div class=\"control-language\">\n<div class=\"dm-buttons\">\n<div class=\"dm-buttons-left\">\n<div class=\"dm-button-snippet red-button\"><\/div>\n<div class=\"dm-button-snippet orange-button\"><\/div>\n<div class=\"dm-button-snippet green-button\"><\/div>\n<\/div>\n<div class=\"dm-buttons-right\"><a><span class=\"dm-copy-text\">Copy Code<\/span><span class=\"dm-copy-confirmed\">Copied<\/span><span class=\"dm-error-message\">Use a different Browser<\/span><\/a><\/div>\n<\/div>\n<pre class=\"no-line-numbers\"><code class=\"no-wrap language-php\">import os, sys, subprocess, importlib.util, functools, dataclasses, time, math\ndef _sh(cmd):\n   subprocess.run(cmd, shell=True, check=False,\n                  stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL)\nprint(\"Installing dependencies ...\")\n_sh(f'{sys.executable} -m pip install -q \"etils[array-types,epy,etree,enp]\" '\n   f'chex flax optax scikit-image')\nREPO_DIR = \"\/content\/jax3d\" if os.path.isdir(\"\/content\") else os.path.abspath(\".\/jax3d\")\nif not os.path.isdir(REPO_DIR):\n   print(\"Cloning google-research\/jax3d ...\")\n   _sh(f\"git clone -q --depth 1 https:\/\/github.com\/google-research\/jax3d.git {REPO_DIR}\")\ndef _load_module_by_path(name, path):\n   \"\"\"Load a single .py file without triggering the parent package __init__.\n   `from jax3d.math import volume_rendering` also works if you run\n   `pip install .` inside the clone, but that pulls in gin\/tfds\/etc.\n   \"\"\"\n   spec = importlib.util.spec_from_file_location(name, path)\n   mod = importlib.util.module_from_spec(spec)\n   sys.modules[name] = mod\n   spec.loader.exec_module(mod)\n   return mod\n_VR_PATH = os.path.join(REPO_DIR, \"jax3d\", \"jax3d\", \"math\", \"volume_rendering.py\")\nif not os.path.exists(_VR_PATH):\n   _VR_PATH = os.path.join(REPO_DIR, \"jax3d\", \"math\", \"volume_rendering.py\")\ntry:\n   j3vr = _load_module_by_path(\"j3d_volume_rendering\", _VR_PATH)\nexcept Exception as e:\n   raise SystemExit(\n       f\"Could not load {_VR_PATH}: {e}n\"\n       \"Try: pip install -U 'etils[array-types,epy,etree,enp]==1.9.4' and re-run.\"\n   )\nimport numpy as np\nimport jax\nimport jax.numpy as jnp\nimport flax.linen as nn\nimport optax\nfrom flax.training import train_state\nimport matplotlib.pyplot as plt\nfrom PIL import Image\nprint(\"jax\", jax.__version__, \"| device:\", jax.devices()[0].device_kind,\n     f\"({jax.devices()[0].platform})\")\nprint(\"jax3d volume_rendering API:\",\n     [n for n in (\"sample_along_rays\", \"volume_rendering\",\n                  \"sample_piecewise_constant_pdf\", \"sample_1d\")\n      if hasattr(j3vr, n)])\n@dataclasses.dataclass\nclass Config:\n   H: int = 64;            W: int = 64\n   n_train_views: int = 24; n_test_views: int = 3\n   cam_radius: float = 3.2; fov_deg: float = 40.0\n   near: float = 1.9;       far: float = 4.7\n   gt_samples: int = 256\n   n_coarse: int = 64;      n_fine: int = 64\n   deg_pos: int = 10;       deg_dir: int = 4\n   width: int = 128;        depth: int = 6;   skip: int = 3\n   batch_rays: int = 2048;  steps: int = 2500\n   lr_init: float = 5e-4;   lr_final: float = 5e-6\n   chunk: int = 4096\n   grid_res: int = 96\ncfg = Config()\nif jax.devices()[0].platform == \"cpu\":\n   print(\"n!! No GPU detected -- switching to a small CPU-friendly config.\")\n   print(\"   (Runtime &gt; Change runtime type &gt; T4 GPU for the full version.)n\")\n   cfg = dataclasses.replace(cfg, H=40, W=40, n_train_views=14, steps=400,\n                             gt_samples=128, n_coarse=32, n_fine=32,\n                             width=64, depth=4, skip=2, batch_rays=1024,\n                             chunk=1600, grid_res=64)\ndef _normalize(v, axis=-1):\n   return v \/ (np.linalg.norm(v, axis=axis, keepdims=True) + 1e-9)\ndef look_at(eye, target=(0., 0., 0.), up=(0., 0., 1.)):\n   \"\"\"OpenGL\/NeRF convention camera-to-world: +x right, +y up, camera looks at -z.\"\"\"\n   eye, target, up = map(lambda a: np.asarray(a, np.float32), (eye, target, up))\n   fwd   = _normalize(target - eye)\n   right = _normalize(np.cross(fwd, up))\n   trueup = np.cross(right, fwd)\n   c2w = np.eye(4, dtype=np.float32)\n   c2w[:3, :3] = np.stack([right, trueup, -fwd], axis=1)\n   c2w[:3, 3] = eye\n   return c2w\ndef orbit_poses(n, radius, elev_lo=18., elev_hi=58., phase=0.0):\n   \"\"\"Golden-angle azimuths + monotone elevations =&gt; well-spread views on a dome.\"\"\"\n   i = np.arange(n, dtype=np.float64) + 0.5\n   az = 2 * np.pi * ((i * 0.6180339887) + phase)\n   elev = np.arcsin(np.linspace(np.sin(np.deg2rad(elev_lo)),\n                                np.sin(np.deg2rad(elev_hi)), n))\n   eyes = np.stack([radius * np.cos(elev) * np.cos(az),\n                    radius * np.cos(elev) * np.sin(az),\n                    radius * np.sin(elev)], axis=-1).astype(np.float32)\n   return np.stack([look_at(e) for e in eyes], axis=0)\ndef rays_from_pose(c2w, H, W, focal):\n   \"\"\"Returns (origins, dirs) of shape [H, W, 3]; dirs are unit-length, so the\n   depths returned by jax3d's sampler are true world-space distances.\"\"\"\n   i, j = np.meshgrid(np.arange(W, dtype=np.float32),\n                      np.arange(H, dtype=np.float32), indexing=\"xy\")\n   cam_dirs = np.stack([(i - W * .5 + .5) \/ focal,\n                        -(j - H * .5 + .5) \/ focal,\n                        -np.ones_like(i)], axis=-1)\n   dirs = _normalize(cam_dirs @ c2w[:3, :3].T)\n   origins = np.broadcast_to(c2w[:3, 3], dirs.shape)\n   return origins.astype(np.float32).copy(), dirs.astype(np.float32)\nFOCAL = 0.5 * cfg.W \/ math.tan(0.5 * math.radians(cfg.fov_deg))\n<\/code><\/pre>\n<\/div>\n<\/div>\n<p class=\"wp-block-paragraph\">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.<\/p>\n<div class=\"dm-code-snippet dark dm-normal-version default no-background-mobile\">\n<div class=\"control-language\">\n<div class=\"dm-buttons\">\n<div class=\"dm-buttons-left\">\n<div class=\"dm-button-snippet red-button\"><\/div>\n<div class=\"dm-button-snippet orange-button\"><\/div>\n<div class=\"dm-button-snippet green-button\"><\/div>\n<\/div>\n<div class=\"dm-buttons-right\"><a><span class=\"dm-copy-text\">Copy Code<\/span><span class=\"dm-copy-confirmed\">Copied<\/span><span class=\"dm-error-message\">Use a different Browser<\/span><\/a><\/div>\n<\/div>\n<pre class=\"no-line-numbers\"><code class=\"no-wrap language-php\">LIGHT = jnp.asarray(_normalize(np.array([0.55, 0.75, 0.85], np.float32)))\n_SPHERES = [\n   (jnp.array([0.34, 0.02, -0.22]), 0.36, jnp.array([0.90, 0.24, 0.22])),\n   (jnp.array([-0.32, 0.28, 0.05]), 0.26, jnp.array([0.25, 0.78, 0.36])),\n   (jnp.array([-0.05, -0.36, 0.24]), 0.22, jnp.array([0.28, 0.40, 0.95])),\n]\ndef _sphere_field(pos, vdir, center, radius, albedo):\n   d = pos - center\n   dist = jnp.linalg.norm(d, axis=-1)\n   n = d \/ (dist[..., None] + 1e-8)\n   sigma = 80.0 * jax.nn.sigmoid((radius - dist) \/ 0.015)\n   v = -vdir\n   refl = 2.0 * jnp.sum(n * v, -1, keepdims=True) * n - v\n   spec = 0.65 * jnp.clip(jnp.sum(refl * LIGHT, -1), 0., 1.) ** 24\n   lamb = 0.35 + 0.65 * jnp.clip(jnp.sum(n * LIGHT, -1), 0., 1.)\n   rgb = jnp.clip(albedo * lamb[..., None] + spec[..., None], 0., 1.)\n   return sigma, rgb\ndef _floor_field(pos):\n   x, y, z = pos[..., 0], pos[..., 1], pos[..., 2]\n   m = (jax.nn.sigmoid((0.06 - jnp.abs(z + 0.62)) \/ 0.008)\n        * jax.nn.sigmoid((0.85 - jnp.abs(x)) \/ 0.01)\n        * jax.nn.sigmoid((0.85 - jnp.abs(y)) \/ 0.01))\n   checker = (jnp.floor(x * 3.0) + jnp.floor(y * 3.0)) % 2.0\n   rgb = jnp.where(checker[..., None] &gt; 0.5,\n                   jnp.array([0.86, 0.86, 0.89]), jnp.array([0.22, 0.25, 0.30]))\n   return 80.0 * m, rgb\ndef gt_field(pos, vdir):\n   \"\"\"pos, vdir: [..., 3] -&gt; (sigma [...], rgb [..., 3]). Density-weighted blend.\"\"\"\n   sig_sum = 0.0\n   col_sum = 0.0\n   for c, r, a in _SPHERES:\n       s, rgb = _sphere_field(pos, vdir, c, r, a)\n       sig_sum = sig_sum + s\n       col_sum = col_sum + s[..., None] * rgb\n   s, rgb = _floor_field(pos)\n   sig_sum = sig_sum + s\n   col_sum = col_sum + s[..., None] * rgb\n   return sig_sum, col_sum \/ (sig_sum[..., None] + 1e-8)\nWHITE_BG = jnp.ones((3,), jnp.float32)\n@jax.jit\ndef render_ground_truth(origins, dirs):\n   \"\"\"Fine-grained volumetric render of the analytic scene -&gt; RGB + depth.\"\"\"\n   depths, positions = j3vr.sample_along_rays(\n       ray_origins=origins, ray_directions=dirs,\n       near=cfg.near, far=cfg.far,\n       sample_count=cfg.gt_samples, deterministic=True)\n   vdir = jnp.broadcast_to(dirs[..., None, :], positions.shape)\n   sigma, rgb = gt_field(positions, vdir)\n   out = j3vr.volume_rendering(\n       sample_values={\"rgb\": rgb}, sample_density=sigma, depths=depths,\n       background_values={\"rgb\": WHITE_BG})\n   return out.ray_values[\"rgb\"], out.ray_depth, out.ray_alpha\ndef build_dataset(poses):\n   O, D, C = [], [], []\n   for c2w in poses:\n       o, d = rays_from_pose(c2w, cfg.H, cfg.W, FOCAL)\n       rgb, _, _ = render_ground_truth(jnp.asarray(o), jnp.asarray(d))\n       O.append(o); D.append(d); C.append(np.asarray(rgb))\n   return (np.stack(O), np.stack(D), np.stack(C))\nprint(\"nRendering the synthetic multi-view dataset ...\")\nt0 = time.time()\ntrain_poses = orbit_poses(cfg.n_train_views, cfg.cam_radius, phase=0.00)\ntest_poses  = orbit_poses(cfg.n_test_views,  cfg.cam_radius, 26., 50., phase=0.41)\ntr_o, tr_d, tr_c = build_dataset(train_poses)\nte_o, te_d, te_c = build_dataset(test_poses)\nprint(f\"  {cfg.n_train_views} train + {cfg.n_test_views} test views \"\n     f\"at {cfg.H}x{cfg.W}  ({time.time()-t0:.1f}s)\")\nk = min(8, cfg.n_train_views)\nfig, axes = plt.subplots(1, k, figsize=(2 * k, 2.3))\nfor a, im, p in zip(axes, tr_c[:k], train_poses[:k]):\n   a.imshow(np.clip(im, 0, 1)); a.axis(\"off\")\n   a.set_title(f\"({p[0,3]:+.1f},{p[1,3]:+.1f},{p[2,3]:+.1f})\", fontsize=7)\nfig.suptitle(\"Training views (ground truth, rendered with jax3d.math.volume_rendering)\",\n            fontsize=11); plt.tight_layout(); plt.show()\nrays_o = jnp.asarray(tr_o.reshape(-1, 3))\nrays_d = jnp.asarray(tr_d.reshape(-1, 3))\nrays_c = jnp.asarray(tr_c.reshape(-1, 3))\nN_RAYS = rays_o.shape[0]\nprint(f\"  ray pool: {N_RAYS:,} rays\")\n<\/code><\/pre>\n<\/div>\n<\/div>\n<p class=\"wp-block-paragraph\">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\u2019s 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.<\/p>\n<div class=\"dm-code-snippet dark dm-normal-version default no-background-mobile\">\n<div class=\"control-language\">\n<div class=\"dm-buttons\">\n<div class=\"dm-buttons-left\">\n<div class=\"dm-button-snippet red-button\"><\/div>\n<div class=\"dm-button-snippet orange-button\"><\/div>\n<div class=\"dm-button-snippet green-button\"><\/div>\n<\/div>\n<div class=\"dm-buttons-right\"><a><span class=\"dm-copy-text\">Copy Code<\/span><span class=\"dm-copy-confirmed\">Copied<\/span><span class=\"dm-error-message\">Use a different Browser<\/span><\/a><\/div>\n<\/div>\n<pre class=\"no-line-numbers\"><code class=\"no-wrap language-php\">def posenc(x, deg):\n   \"\"\"NeRF sinusoidal encoding, with the raw input concatenated.\"\"\"\n   if deg == 0:\n       return x\n   scales = 2.0 ** jnp.arange(deg, dtype=x.dtype)\n   xb = (x[..., None, :] * scales[:, None]).reshape(*x.shape[:-1], -1)\n   return jnp.concatenate([x, jnp.sin(xb), jnp.cos(xb)], axis=-1)\nclass NeRFMLP(nn.Module):\n   width: int; depth: int; skip: int; deg_pos: int; deg_dir: int\n   @nn.compact\n   def __call__(self, pos, dirs):\n       inp = posenc(pos, self.deg_pos)\n       x = inp\n       for i in range(self.depth):\n           x = nn.relu(nn.Dense(self.width)(x))\n           if i == self.skip:\n               x = jnp.concatenate([x, inp], axis=-1)\n       sigma = nn.softplus(nn.Dense(1)(x)[..., 0] - 1.0)\n       h = jnp.concatenate([nn.Dense(self.width)(x), posenc(dirs, self.deg_dir)], -1)\n       rgb = nn.sigmoid(nn.Dense(3)(nn.relu(nn.Dense(self.width \/\/ 2)(h))))\n       return sigma, rgb\nmodel = NeRFMLP(cfg.width, cfg.depth, cfg.skip, cfg.deg_pos, cfg.deg_dir)\n<\/code><\/pre>\n<\/div>\n<\/div>\n<p class=\"wp-block-paragraph\">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.<\/p>\n<div class=\"dm-code-snippet dark dm-normal-version default no-background-mobile\">\n<div class=\"control-language\">\n<div class=\"dm-buttons\">\n<div class=\"dm-buttons-left\">\n<div class=\"dm-button-snippet red-button\"><\/div>\n<div class=\"dm-button-snippet orange-button\"><\/div>\n<div class=\"dm-button-snippet green-button\"><\/div>\n<\/div>\n<div class=\"dm-buttons-right\"><a><span class=\"dm-copy-text\">Copy Code<\/span><span class=\"dm-copy-confirmed\">Copied<\/span><span class=\"dm-error-message\">Use a different Browser<\/span><\/a><\/div>\n<\/div>\n<pre class=\"no-line-numbers\"><code class=\"no-wrap language-php\">def render_rays(params, origins, dirs, rng, deterministic):\n   \"\"\"Coarse pass -&gt; importance-resample -&gt; fine pass. All sampling and\n   compositing comes from jax3d.math.volume_rendering.\"\"\"\n   rng_c, rng_f = jax.random.split(rng)\n   depths_c, pos_c = j3vr.sample_along_rays(\n       ray_origins=origins, ray_directions=dirs,\n       near=cfg.near, far=cfg.far, sample_count=cfg.n_coarse,\n       deterministic=deterministic, rng=rng_c)\n   dirs_c = jnp.broadcast_to(dirs[:, None, :], pos_c.shape)\n   sigma_c, rgb_c = model.apply(params[\"coarse\"], pos_c, dirs_c)\n   out_c = j3vr.volume_rendering(\n       sample_values={\"rgb\": rgb_c}, sample_density=sigma_c, depths=depths_c,\n       background_values={\"rgb\": WHITE_BG})\n   mid = 0.5 * (depths_c[..., 1:] + depths_c[..., :-1])\n   bin_edges = jnp.concatenate([depths_c[..., :1], mid, depths_c[..., -1:]], -1)\n   t_fine = j3vr.sample_piecewise_constant_pdf(\n       bin_edges=bin_edges, weights=out_c.sample_weights,\n       sample_count=cfg.n_fine, deterministic=deterministic, rng=rng_f)\n   t_fine = jax.lax.stop_gradient(t_fine)\n   depths_f = jnp.sort(jnp.concatenate([depths_c, t_fine], -1), axis=-1)\n   pos_f = origins[:, None, :] + depths_f[..., None] * dirs[:, None, :]\n   dirs_f = jnp.broadcast_to(dirs[:, None, :], pos_f.shape)\n   sigma_f, rgb_f = model.apply(params[\"fine\"], pos_f, dirs_f)\n   out_f = j3vr.volume_rendering(\n       sample_values={\"rgb\": rgb_f}, sample_density=sigma_f, depths=depths_f,\n       background_values={\"rgb\": WHITE_BG})\n   aux = {\"depths_c\": depths_c, \"weights_c\": out_c.sample_weights, \"t_fine\": t_fine}\n   return out_c, out_f, aux\ndef mse_to_psnr(x):\n   return -10.0 * jnp.log10(jnp.maximum(x, 1e-10))\n<\/code><\/pre>\n<\/div>\n<\/div>\n<p class=\"wp-block-paragraph\">We implement the core hierarchical renderer by first sampling coarse points along each ray and compositing their densities and colors through JAX3D\u2019s volume-rendering operator. We convert the resulting coarse rendering weights into a piecewise-constant probability distribution and importance-sample additional fine points around high-contribution regions. We combine and sort the coarse and fine samples before performing the final fine-network rendering, while stopping gradients through the sampling operation.<\/p>\n<div class=\"dm-code-snippet dark dm-normal-version default no-background-mobile\">\n<div class=\"control-language\">\n<div class=\"dm-buttons\">\n<div class=\"dm-buttons-left\">\n<div class=\"dm-button-snippet red-button\"><\/div>\n<div class=\"dm-button-snippet orange-button\"><\/div>\n<div class=\"dm-button-snippet green-button\"><\/div>\n<\/div>\n<div class=\"dm-buttons-right\"><a><span class=\"dm-copy-text\">Copy Code<\/span><span class=\"dm-copy-confirmed\">Copied<\/span><span class=\"dm-error-message\">Use a different Browser<\/span><\/a><\/div>\n<\/div>\n<pre class=\"no-line-numbers\"><code class=\"no-wrap language-php\">key = jax.random.PRNGKey(0)\nkey, k1, k2 = jax.random.split(key, 3)\ndummy_p = jnp.zeros((1, 1, 3)); dummy_d = jnp.zeros((1, 1, 3))\nparams = {\"coarse\": model.init(k1, dummy_p, dummy_d),\n         \"fine\":   model.init(k2, dummy_p, dummy_d)}\nn_params = sum(x.size for x in jax.tree.leaves(params))\nprint(f\"nModel: {n_params\/1e6:.2f}M parameters (coarse + fine networks)\")\nschedule = optax.exponential_decay(cfg.lr_init, cfg.steps,\n                                  cfg.lr_final \/ cfg.lr_init)\ntx = optax.chain(optax.clip_by_global_norm(1.0), optax.adam(schedule))\nstate = train_state.TrainState.create(apply_fn=model.apply, params=params, tx=tx)\n@jax.jit\ndef train_step(state, o, d, target, rng):\n   def loss_fn(p):\n       out_c, out_f, _ = render_rays(p, o, d, rng, deterministic=False)\n       l_c = jnp.mean((out_c.ray_values[\"rgb\"] - target) ** 2)\n       l_f = jnp.mean((out_f.ray_values[\"rgb\"] - target) ** 2)\n       return l_c + l_f, l_f\n   (loss, l_fine), grads = jax.value_and_grad(loss_fn, has_aux=True)(state.params)\n   return state.apply_gradients(grads=grads), loss, l_fine\nprint(f\"Training {cfg.steps} steps x {cfg.batch_rays} rays \"\n     f\"({cfg.n_coarse} coarse + {cfg.n_coarse + cfg.n_fine} fine samples\/ray) ...\")\nhistory = []\nt0 = time.time()\nfor step in range(1, cfg.steps + 1):\n   key, k_idx, k_render = jax.random.split(key, 3)\n   idx = jax.random.randint(k_idx, (cfg.batch_rays,), 0, N_RAYS)\n   state, loss, l_fine = train_step(state, rays_o[idx], rays_d[idx],\n                                    rays_c[idx], k_render)\n   if step % 25 == 0 or step == 1:\n       history.append((step, float(mse_to_psnr(l_fine))))\n   if step % max(1, cfg.steps \/\/ 10) == 0 or step == 1:\n       print(f\"  step {step:5d}\/{cfg.steps} | loss {float(loss):.5f} \"\n             f\"| train PSNR {float(mse_to_psnr(l_fine)):5.2f} dB \"\n             f\"| {time.time()-t0:6.1f}s\")\nprint(f\"Done in {time.time()-t0:.1f}s\")\n<\/code><\/pre>\n<\/div>\n<\/div>\n<p class=\"wp-block-paragraph\">We initialize independent coarse and fine NeRF networks and optimize them jointly with Adam, using exponential learning-rate decay and global gradient clipping. We supervise both rendering stages against ground-truth ray colors, encouraging the coarse network to learn useful sampling distributions while improving the final fine reconstruction. We run the training step with JAX JIT compilation and monitor the fine-network PSNR throughout optimization.<\/p>\n<div class=\"dm-code-snippet dark dm-normal-version default no-background-mobile\">\n<div class=\"control-language\">\n<div class=\"dm-buttons\">\n<div class=\"dm-buttons-left\">\n<div class=\"dm-button-snippet red-button\"><\/div>\n<div class=\"dm-button-snippet orange-button\"><\/div>\n<div class=\"dm-button-snippet green-button\"><\/div>\n<\/div>\n<div class=\"dm-buttons-right\"><a><span class=\"dm-copy-text\">Copy Code<\/span><span class=\"dm-copy-confirmed\">Copied<\/span><span class=\"dm-error-message\">Use a different Browser<\/span><\/a><\/div>\n<\/div>\n<pre class=\"no-line-numbers\"><code class=\"no-wrap language-php\">@jax.jit\ndef render_chunk(params, o, d, rng):\n   _, out_f, aux = render_rays(params, o, d, rng, deterministic=True)\n   return out_f.ray_values[\"rgb\"], out_f.ray_depth, out_f.ray_alpha, aux\ndef render_image(params, origins, dirs, rng):\n   \"\"\"Chunked full-image render with padding, so only one shape gets compiled.\"\"\"\n   o = jnp.asarray(origins.reshape(-1, 3)); d = jnp.asarray(dirs.reshape(-1, 3))\n   R = o.shape[0]; rgb, dep, alp = [], [], []\n   for i in range(0, R, cfg.chunk):\n       oc, dc = o[i:i + cfg.chunk], d[i:i + cfg.chunk]\n       pad = cfg.chunk - oc.shape[0]\n       if pad:\n           oc = jnp.concatenate([oc, jnp.tile(oc[-1:], (pad, 1))], 0)\n           dc = jnp.concatenate([dc, jnp.tile(dc[-1:], (pad, 1))], 0)\n       c, dp, a, _ = render_chunk(params, oc, dc, rng)\n       n = cfg.chunk - pad\n       rgb.append(c[:n]); dep.append(dp[:n]); alp.append(a[:n])\n   s = (cfg.H, cfg.W)\n   return (np.asarray(jnp.concatenate(rgb)).reshape(*s, 3),\n           np.asarray(jnp.concatenate(dep)).reshape(*s),\n           np.asarray(jnp.concatenate(alp)).reshape(*s))\nh = np.array(history)\nplt.figure(figsize=(6, 3))\nplt.plot(h[:, 0], h[:, 1], lw=1.6)\nplt.xlabel(\"step\"); plt.ylabel(\"train PSNR (dB)\")\nplt.title(\"Fine-network training PSNR\"); plt.grid(alpha=.3)\nplt.tight_layout(); plt.show()\nprint(\"nRendering held-out test views ...\")\nkey, k_eval = jax.random.split(key)\npsnrs = []\nfig, axes = plt.subplots(cfg.n_test_views, 4,\n                        figsize=(11, 2.7 * cfg.n_test_views), squeeze=False)\nfor v in range(cfg.n_test_views):\n   pred, depth, alpha = render_image(state.params, te_o[v], te_d[v], k_eval)\n   p = float(mse_to_psnr(np.mean((pred - te_c[v]) ** 2))); psnrs.append(p)\n   depth_vis = depth + (1.0 - alpha) * cfg.far\n   for a, (im, ttl, kw) in zip(axes[v], [\n           (np.clip(te_c[v], 0, 1), \"ground truth\", {}),\n           (np.clip(pred, 0, 1), f\"NeRF  ({p:.2f} dB)\", {}),\n           (depth_vis, \"depth (ray_depth)\", dict(cmap=\"turbo\",\n                                                 vmin=cfg.near, vmax=cfg.far)),\n           (alpha, \"opacity (ray_alpha)\", dict(cmap=\"gray\", vmin=0, vmax=1))]):\n       a.imshow(im, **kw); a.set_title(ttl, fontsize=9); a.axis(\"off\")\nplt.suptitle(f\"Novel-view synthesis   |   mean PSNR = {np.mean(psnrs):.2f} dB\",\n            fontsize=12)\nplt.tight_layout(); plt.show()\nprint(f\"  mean held-out PSNR: {np.mean(psnrs):.2f} dB\")\ncy, cx = cfg.H \/\/ 2, cfg.W \/\/ 2\no1 = jnp.asarray(te_o[0][cy, cx])[None]; d1 = jnp.asarray(te_d[0][cy, cx])[None]\no1 = jnp.tile(o1, (cfg.chunk, 1)); d1 = jnp.tile(d1, (cfg.chunk, 1))\n_, _, _, aux = render_chunk(state.params, o1, d1, k_eval)\ndc = np.asarray(aux[\"depths_c\"][0]); wc = np.asarray(aux[\"weights_c\"][0])\ntf = np.asarray(aux[\"t_fine\"][0])\nfig, ax = plt.subplots(figsize=(8, 3))\nax.bar(dc, wc, width=(cfg.far - cfg.near) \/ cfg.n_coarse * .9,\n      alpha=.55, label=\"coarse weights (the PDF)\")\nax.plot(tf, np.full_like(tf, wc.max() * .06), \"|\", ms=16, color=\"crimson\",\n       label=\"fine samples (sample_piecewise_constant_pdf)\")\nax.set_xlabel(\"depth along ray\"); ax.set_ylabel(\"weight\")\nax.set_title(\"Importance resampling concentrates samples on the surface\")\nax.legend(fontsize=8); plt.tight_layout(); plt.show()\nprint(\"nRendering 360-degree orbit ...\")\nn_frames = 24 if jax.devices()[0].platform != \"cpu\" else 8\nframes = []\nfor t in range(n_frames):\n   az = 2 * np.pi * t \/ n_frames; el = np.deg2rad(32.0)\n   eye = cfg.cam_radius * np.array([np.cos(el) * np.cos(az),\n                                    np.cos(el) * np.sin(az), np.sin(el)])\n   o, d = rays_from_pose(look_at(eye), cfg.H, cfg.W, FOCAL)\n   rgb, _, _ = render_image(state.params, o, d, k_eval)\n   frames.append((np.clip(rgb, 0, 1) * 255).astype(np.uint8))\ngif_path = os.path.join(os.getcwd(), \"nerf_orbit.gif\")\npil = [Image.fromarray(f).resize((cfg.W * 3, cfg.H * 3), Image.NEAREST) for f in frames]\npil[0].save(gif_path, save_all=True, append_images=pil[1:], duration=90, loop=0)\ntry:\n   from IPython.display import Image as IPImage, display\n   display(IPImage(filename=gif_path))\nexcept Exception:\n   pass\nprint(\"  saved\", gif_path)\nprint(\"nExtracting isosurface from the learned density field ...\")\ntry:\n   from skimage import measure\n   g = np.linspace(-1.0, 1.0, cfg.grid_res, dtype=np.float32)\n   X, Y, Z = np.meshgrid(g, g, g, indexing=\"ij\")\n   pts = np.stack([X, Y, Z], -1).reshape(-1, 3)\n   @jax.jit\n   def density_at(p):\n       s, _ = model.apply(state.params[\"fine\"], p, jnp.zeros_like(p))\n       return s\n   vol = np.concatenate([np.asarray(density_at(jnp.asarray(pts[i:i + 65536])))\n                         for i in range(0, pts.shape[0], 65536)])\n   vol = vol.reshape(cfg.grid_res, cfg.grid_res, cfg.grid_res)\n   step = (cfg.far - cfg.near) \/ (cfg.n_coarse + cfg.n_fine)\n   level = float(-np.log(0.5) \/ step)\n   if not (vol.min() &lt; level &lt; vol.max()):\n       level = float(np.percentile(vol, 99.0))\n   verts, faces, _, _ = measure.marching_cubes(vol, level=level)\n   verts = -1.0 + verts * (2.0 \/ (cfg.grid_res - 1))\n   fig = plt.figure(figsize=(6, 6)); ax = fig.add_subplot(111, projection=\"3d\")\n   ax.plot_trisurf(verts[:, 0], verts[:, 1], verts[:, 2], triangles=faces,\n                   cmap=\"viridis\", lw=0.0, antialiased=False, alpha=.95)\n   ax.set_box_aspect((1, 1, 1))\n   ax.set_xlim(-1, 1); ax.set_ylim(-1, 1); ax.set_zlim(-1, 1)\n   ax.view_init(elev=24, azim=-58)\n   ax.set_title(f\"Marching cubes on learned density  (sigma = {level:.1f}, \"\n                f\"{len(faces):,} faces)\", fontsize=10)\n   plt.tight_layout(); plt.show()\nexcept Exception as e:\n   print(\"  isosurface step skipped:\", e)\nprint(\"n\" + \"=\" * 70)\nprint(f\"FINAL held-out PSNR: {np.mean(psnrs):.2f} dB   ({n_params\/1e6:.2f}M params, \"\n     f\"{cfg.steps} steps)\")\nprint(\"jax3d functions exercised: sample_along_rays, volume_rendering, \"\n     \"sample_piecewise_constant_pdf\")\nprint(\"=\" * 70)\n<\/code><\/pre>\n<\/div>\n<\/div>\n<p class=\"wp-block-paragraph\">We evaluate the trained representation through chunked novel-view rendering and measure reconstruction quality with held-out PSNR, along with depth and opacity maps. We visualize how hierarchical sampling concentrates fine samples around important surfaces, then generate a 360-degree orbit GIF to inspect the learned radiance field from multiple viewpoints. We finally query the learned density on a 3D grid and apply marching cubes to extract an approximate geometric isosurface.<\/p>\n<p class=\"wp-block-paragraph\">In conclusion, we demonstrated the complete inverse-rendering pipeline by learning a continuous density and radiance field from synthetic multi-view observations and reconstructing it through hierarchical volume rendering. We used the coarse network to identify informative regions along each ray and the fine network to concentrate additional samples around high-contribution surfaces. At the same time, view-direction encoding allows us to model view-dependent appearance. In the final evaluation stages, we measured novel-view reconstruction quality with PSNR, inspected learned depth and opacity, visualized importance-sampling behavior, generated a 360-degree orbit, and extracted an approximate learned geometry with marching cubes. Overall, we showed how the mathematical components of jax3d integrate with modern JAX-based neural-network training to form a compact yet technically complete NeRF reconstruction system.<\/p>\n<p class=\"wp-block-paragraph\">\n<hr class=\"wp-block-separator has-alpha-channel-opacity\" \/>\n<\/p><p class=\"wp-block-paragraph\">\n<\/p><p class=\"wp-block-paragraph\">Check out <strong><a href=\"https:\/\/github.com\/MARKTECHPOST-AI-MEDIA-INC\/AI-Agents-Projects-Tutorials\/blob\/main\/Computer%20Vision\/jax3d_hierarchical_nerf_tutorial_Marktechpost.ipynb\" target=\"_blank\" rel=\"noreferrer noopener\">the\u00a0FULL CODES here<\/a><\/strong>. All credit goes to the researcher of this project. Also,\u00a0feel free to follow us on\u00a0<strong><a href=\"https:\/\/x.com\/intent\/follow?screen_name=marktechpost\" target=\"_blank\" rel=\"noopener\"><mark>Twitter<\/mark><\/a><\/strong>\u00a0and don\u2019t forget to join our\u00a0<strong><a href=\"https:\/\/www.reddit.com\/r\/machinelearningnews\/\" target=\"_blank\" rel=\"noopener\">150k+ML SubReddit<\/a><\/strong>\u00a0and Subscribe to\u00a0<strong><a href=\"https:\/\/magic.beehiiv.com\/v1\/f5e63dd4-5653-4f09-83e2-321a8b1ba526?email=%7B%7Bemail%7D%7D\" target=\"_blank\" rel=\"noopener\">our Newsletter<\/a><\/strong>. Wait! are you on telegram?\u00a0<strong><a href=\"https:\/\/t.me\/machinelearningresearchnews\" target=\"_blank\" rel=\"noopener\">now you can join us on telegram as well.<\/a><\/strong><\/p>\n<p class=\"wp-block-paragraph\">Need to partner with us for promoting your GitHub Repo OR Hugging Face Page OR Product Release OR Webinar etc.?\u00a0<strong><a href=\"https:\/\/forms.gle\/wbash1wF6efRj8G58\" target=\"_blank\" rel=\"noopener\"><mark>Connect with us<\/mark><\/a><\/strong><\/p>\n<p>The post <a href=\"https:\/\/www.marktechpost.com\/2026\/09\/13\/hierarchical-nerf-with-jax3d-for-volumetric-rendering-novel-view-synthesis-and-3d-reconstruction\/\">Hierarchical NeRF with JAX3D for Volumetric Rendering, Novel-View Synthesis, and 3D Reconstruction<\/a> appeared first on <a href=\"https:\/\/www.marktechpost.com\/\">MarkTechPost<\/a>.<\/p>","protected":false},"excerpt":{"rendered":"<p>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(&#8220;Installing dependencies &#8230;&#8221;) _sh(f'{sys.executable} -m pip install -q &#8220;etils[array-types,epy,etree,enp]&#8221; &#8216; f&#8217;chex flax optax scikit-image&#8217;) REPO_DIR = &#8220;\/content\/jax3d&#8221; if os.path.isdir(&#8220;\/content&#8221;) else os.path.abspath(&#8220;.\/jax3d&#8221;) if not os.path.isdir(REPO_DIR): print(&#8220;Cloning google-research\/jax3d &#8230;&#8221;) _sh(f&#8221;git clone -q &#8211;depth 1 https:\/\/github.com\/google-research\/jax3d.git {REPO_DIR}&#8221;) def _load_module_by_path(name, path): &#8220;&#8221;&#8221;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. &#8220;&#8221;&#8221; 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, &#8220;jax3d&#8221;, &#8220;jax3d&#8221;, &#8220;math&#8221;, &#8220;volume_rendering.py&#8221;) if not os.path.exists(_VR_PATH): _VR_PATH = os.path.join(REPO_DIR, &#8220;jax3d&#8221;, &#8220;math&#8221;, &#8220;volume_rendering.py&#8221;) try: j3vr = _load_module_by_path(&#8220;j3d_volume_rendering&#8221;, _VR_PATH) except Exception as e: raise SystemExit( f&#8221;Could not load {_VR_PATH}: {e}n&#8221; &#8220;Try: pip install -U &#8216;etils[array-types,epy,etree,enp]==1.9.4&#8217; and re-run.&#8221; ) 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(&#8220;jax&#8221;, jax.__version__, &#8220;| device:&#8221;, jax.devices()[0].device_kind, f&#8221;({jax.devices()[0].platform})&#8221;) print(&#8220;jax3d volume_rendering API:&#8221;, [n for n in (&#8220;sample_along_rays&#8221;, &#8220;volume_rendering&#8221;, &#8220;sample_piecewise_constant_pdf&#8221;, &#8220;sample_1d&#8221;) 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 == &#8220;cpu&#8221;: print(&#8220;n!! No GPU detected &#8212; switching to a small CPU-friendly config.&#8221;) print(&#8221; (Runtime &gt; Change runtime type &gt; T4 GPU for the full version.)n&#8221;) 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.)): &#8220;&#8221;&#8221;OpenGL\/NeRF convention camera-to-world: +x right, +y up, camera looks at -z.&#8221;&#8221;&#8221; eye, target, up = map(lambda a: np.asarray(a, np.float32), (eye, target, up)) fwd = _normalize(target &#8211; 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): &#8220;&#8221;&#8221;Golden-angle azimuths + monotone elevations =&gt; well-spread views on a dome.&#8221;&#8221;&#8221; 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): &#8220;&#8221;&#8221;Returns (origins, dirs) of shape [H, W, 3]; dirs are unit-length, so the depths returned by jax3d&#8217;s sampler are true world-space distances.&#8221;&#8221;&#8221; i, j = np.meshgrid(np.arange(W, dtype=np.float32), np.arange(H, dtype=np.float32), indexing=&#8221;xy&#8221;) cam_dirs = np.stack([(i &#8211; W * .5 + .5) \/ focal, -(j &#8211; 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 &#8211; center dist = jnp.linalg.norm(d, axis=-1) n = d \/ (dist[&#8230;, None] + 1e-8) sigma = 80.0 * jax.nn.sigmoid((radius &#8211; dist) \/ 0.015) v = -vdir refl = 2.0 * jnp.sum(n * v, -1, keepdims=True) * n &#8211; 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[&#8230;, None] + spec[&#8230;, None], 0., 1.) return sigma, rgb def _floor_field(pos): x, y, z = pos[&#8230;, 0], pos[&#8230;, 1], pos[&#8230;, 2] m = (jax.nn.sigmoid((0.06 &#8211; jnp.abs(z + 0.62)) \/ 0.008) * jax.nn.sigmoid((0.85 &#8211; jnp.abs(x)) \/ 0.01) * jax.nn.sigmoid((0.85 &#8211; jnp.abs(y)) \/ 0.01)) checker = (jnp.floor(x * 3.0) + jnp.floor(y * 3.0)) % 2.0 rgb = jnp.where(checker[&#8230;, None] &gt; 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): &#8220;&#8221;&#8221;pos, vdir: [&#8230;, 3] -&gt; (sigma [&#8230;], rgb [&#8230;, 3]). Density-weighted blend.&#8221;&#8221;&#8221; 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[&#8230;, None] * rgb s, rgb = _floor_field(pos) sig_sum = sig_sum + s col_sum = col_sum + s[&#8230;, None] * rgb return sig_sum, col_sum \/ (sig_sum[&#8230;, None] + 1e-8) WHITE_BG = jnp.ones((3,), jnp.float32) @jax.jit def render_ground_truth(origins, dirs): &#8220;&#8221;&#8221;Fine-grained volumetric render of the analytic scene -&gt; RGB + depth.&#8221;&#8221;&#8221; 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[&#8230;, None, :],<\/p>","protected":false},"author":2,"featured_media":0,"comment_status":"open","ping_status":"open","sticky":false,"template":"","format":"standard","meta":{"_acf_changed":false,"pmpro_default_level":"","site-sidebar-layout":"default","site-content-layout":"","ast-site-content-layout":"","site-content-style":"default","site-sidebar-style":"default","ast-global-header-display":"","ast-banner-title-visibility":"","ast-main-header-display":"","ast-hfb-above-header-display":"","ast-hfb-below-header-display":"","ast-hfb-mobile-header-display":"","site-post-title":"","ast-breadcrumbs-content":"","ast-featured-img":"","footer-sml-layout":"","theme-transparent-header-meta":"","adv-header-id-meta":"","stick-header-meta":"","header-above-stick-meta":"","header-main-stick-meta":"","header-below-stick-meta":"","astra-migrate-meta-layouts":"default","ast-page-background-enabled":"default","ast-page-background-meta":{"desktop":{"background-color":"var(--ast-global-color-4)","background-image":"","background-repeat":"repeat","background-position":"center center","background-size":"auto","background-attachment":"scroll","background-type":"","background-media":"","overlay-type":"","overlay-color":"","overlay-opacity":"","overlay-gradient":""},"tablet":{"background-color":"","background-image":"","background-repeat":"repeat","background-position":"center center","background-size":"auto","background-attachment":"scroll","background-type":"","background-media":"","overlay-type":"","overlay-color":"","overlay-opacity":"","overlay-gradient":""},"mobile":{"background-color":"","background-image":"","background-repeat":"repeat","background-position":"center center","background-size":"auto","background-attachment":"scroll","background-type":"","background-media":"","overlay-type":"","overlay-color":"","overlay-opacity":"","overlay-gradient":""}},"ast-content-background-meta":{"desktop":{"background-color":"var(--ast-global-color-5)","background-image":"","background-repeat":"repeat","background-position":"center center","background-size":"auto","background-attachment":"scroll","background-type":"","background-media":"","overlay-type":"","overlay-color":"","overlay-opacity":"","overlay-gradient":""},"tablet":{"background-color":"var(--ast-global-color-5)","background-image":"","background-repeat":"repeat","background-position":"center center","background-size":"auto","background-attachment":"scroll","background-type":"","background-media":"","overlay-type":"","overlay-color":"","overlay-opacity":"","overlay-gradient":""},"mobile":{"background-color":"var(--ast-global-color-5)","background-image":"","background-repeat":"repeat","background-position":"center center","background-size":"auto","background-attachment":"scroll","background-type":"","background-media":"","overlay-type":"","overlay-color":"","overlay-opacity":"","overlay-gradient":""}},"_pvb_checkbox_block_on_post":false,"footnotes":""},"categories":[52,5,7,1],"tags":[],"class_list":["post-117670","post","type-post","status-publish","format-standard","hentry","category-ai-club","category-committee","category-news","category-uncategorized","pmpro-has-access"],"acf":[],"yoast_head":"<!-- This site is optimized with the Yoast SEO plugin v25.3 - https:\/\/yoast.com\/wordpress\/plugins\/seo\/ -->\n<title>Hierarchical NeRF with JAX3D for Volumetric Rendering, Novel-View Synthesis, and 3D Reconstruction - YouZum<\/title>\n<meta name=\"description\" content=\"\u0e01\u0e34\u0e08\u0e01\u0e23\u0e23\u0e21\u0e40\u0e01\u0e35\u0e48\u0e22\u0e27\u0e01\u0e31\u0e1a\u0e42\u0e14\u0e23\u0e19\" \/>\n<meta name=\"robots\" content=\"index, follow, max-snippet:-1, max-image-preview:large, max-video-preview:-1\" \/>\n<link rel=\"canonical\" href=\"https:\/\/youzum.net\/ja\/hierarchical-nerf-with-jax3d-for-volumetric-rendering-novel-view-synthesis-and-3d-reconstruction\/\" \/>\n<meta property=\"og:locale\" content=\"ja_JP\" \/>\n<meta property=\"og:type\" content=\"article\" \/>\n<meta property=\"og:title\" content=\"Hierarchical NeRF with JAX3D for Volumetric Rendering, Novel-View Synthesis, and 3D Reconstruction - YouZum\" \/>\n<meta property=\"og:description\" content=\"\u0e01\u0e34\u0e08\u0e01\u0e23\u0e23\u0e21\u0e40\u0e01\u0e35\u0e48\u0e22\u0e27\u0e01\u0e31\u0e1a\u0e42\u0e14\u0e23\u0e19\" \/>\n<meta property=\"og:url\" content=\"https:\/\/youzum.net\/ja\/hierarchical-nerf-with-jax3d-for-volumetric-rendering-novel-view-synthesis-and-3d-reconstruction\/\" \/>\n<meta property=\"og:site_name\" content=\"YouZum\" \/>\n<meta property=\"article:publisher\" content=\"https:\/\/www.facebook.com\/DroneAssociationTH\/\" \/>\n<meta property=\"article:published_time\" content=\"2026-09-14T01:45:47+00:00\" \/>\n<meta name=\"author\" content=\"admin NU\" \/>\n<meta name=\"twitter:card\" content=\"summary_large_image\" \/>\n<meta name=\"twitter:label1\" content=\"\u57f7\u7b46\u8005\" \/>\n\t<meta name=\"twitter:data1\" content=\"admin NU\" \/>\n\t<meta name=\"twitter:label2\" content=\"\u63a8\u5b9a\u8aad\u307f\u53d6\u308a\u6642\u9593\" \/>\n\t<meta name=\"twitter:data2\" content=\"18\u5206\" \/>\n<script type=\"application\/ld+json\" class=\"yoast-schema-graph\">{\"@context\":\"https:\/\/schema.org\",\"@graph\":[{\"@type\":\"Article\",\"@id\":\"https:\/\/youzum.net\/hierarchical-nerf-with-jax3d-for-volumetric-rendering-novel-view-synthesis-and-3d-reconstruction\/#article\",\"isPartOf\":{\"@id\":\"https:\/\/youzum.net\/hierarchical-nerf-with-jax3d-for-volumetric-rendering-novel-view-synthesis-and-3d-reconstruction\/\"},\"author\":{\"name\":\"admin NU\",\"@id\":\"https:\/\/yousum.gpucore.co\/#\/schema\/person\/97fa48242daf3908e4d9a5f26f4a059c\"},\"headline\":\"Hierarchical NeRF with JAX3D for Volumetric Rendering, Novel-View Synthesis, and 3D Reconstruction\",\"datePublished\":\"2026-09-14T01:45:47+00:00\",\"mainEntityOfPage\":{\"@id\":\"https:\/\/youzum.net\/hierarchical-nerf-with-jax3d-for-volumetric-rendering-novel-view-synthesis-and-3d-reconstruction\/\"},\"wordCount\":764,\"commentCount\":0,\"publisher\":{\"@id\":\"https:\/\/yousum.gpucore.co\/#organization\"},\"articleSection\":[\"AI\",\"Committee\",\"News\",\"Uncategorized\"],\"inLanguage\":\"ja\",\"potentialAction\":[{\"@type\":\"CommentAction\",\"name\":\"Comment\",\"target\":[\"https:\/\/youzum.net\/hierarchical-nerf-with-jax3d-for-volumetric-rendering-novel-view-synthesis-and-3d-reconstruction\/#respond\"]}]},{\"@type\":\"WebPage\",\"@id\":\"https:\/\/youzum.net\/hierarchical-nerf-with-jax3d-for-volumetric-rendering-novel-view-synthesis-and-3d-reconstruction\/\",\"url\":\"https:\/\/youzum.net\/hierarchical-nerf-with-jax3d-for-volumetric-rendering-novel-view-synthesis-and-3d-reconstruction\/\",\"name\":\"Hierarchical NeRF with JAX3D for Volumetric Rendering, Novel-View Synthesis, and 3D Reconstruction - YouZum\",\"isPartOf\":{\"@id\":\"https:\/\/yousum.gpucore.co\/#website\"},\"datePublished\":\"2026-09-14T01:45:47+00:00\",\"description\":\"\u0e01\u0e34\u0e08\u0e01\u0e23\u0e23\u0e21\u0e40\u0e01\u0e35\u0e48\u0e22\u0e27\u0e01\u0e31\u0e1a\u0e42\u0e14\u0e23\u0e19\",\"breadcrumb\":{\"@id\":\"https:\/\/youzum.net\/hierarchical-nerf-with-jax3d-for-volumetric-rendering-novel-view-synthesis-and-3d-reconstruction\/#breadcrumb\"},\"inLanguage\":\"ja\",\"potentialAction\":[{\"@type\":\"ReadAction\",\"target\":[\"https:\/\/youzum.net\/hierarchical-nerf-with-jax3d-for-volumetric-rendering-novel-view-synthesis-and-3d-reconstruction\/\"]}]},{\"@type\":\"BreadcrumbList\",\"@id\":\"https:\/\/youzum.net\/hierarchical-nerf-with-jax3d-for-volumetric-rendering-novel-view-synthesis-and-3d-reconstruction\/#breadcrumb\",\"itemListElement\":[{\"@type\":\"ListItem\",\"position\":1,\"name\":\"Home\",\"item\":\"https:\/\/youzum.net\/\"},{\"@type\":\"ListItem\",\"position\":2,\"name\":\"Hierarchical NeRF with JAX3D for Volumetric Rendering, Novel-View Synthesis, and 3D Reconstruction\"}]},{\"@type\":\"WebSite\",\"@id\":\"https:\/\/yousum.gpucore.co\/#website\",\"url\":\"https:\/\/yousum.gpucore.co\/\",\"name\":\"YouSum\",\"description\":\"\",\"publisher\":{\"@id\":\"https:\/\/yousum.gpucore.co\/#organization\"},\"potentialAction\":[{\"@type\":\"SearchAction\",\"target\":{\"@type\":\"EntryPoint\",\"urlTemplate\":\"https:\/\/yousum.gpucore.co\/?s={search_term_string}\"},\"query-input\":{\"@type\":\"PropertyValueSpecification\",\"valueRequired\":true,\"valueName\":\"search_term_string\"}}],\"inLanguage\":\"ja\"},{\"@type\":\"Organization\",\"@id\":\"https:\/\/yousum.gpucore.co\/#organization\",\"name\":\"Drone Association Thailand\",\"url\":\"https:\/\/yousum.gpucore.co\/\",\"logo\":{\"@type\":\"ImageObject\",\"inLanguage\":\"ja\",\"@id\":\"https:\/\/yousum.gpucore.co\/#\/schema\/logo\/image\/\",\"url\":\"https:\/\/youzum.net\/wp-content\/uploads\/2024\/11\/tranparent-logo.png\",\"contentUrl\":\"https:\/\/youzum.net\/wp-content\/uploads\/2024\/11\/tranparent-logo.png\",\"width\":300,\"height\":300,\"caption\":\"Drone Association Thailand\"},\"image\":{\"@id\":\"https:\/\/yousum.gpucore.co\/#\/schema\/logo\/image\/\"},\"sameAs\":[\"https:\/\/www.facebook.com\/DroneAssociationTH\/\"]},{\"@type\":\"Person\",\"@id\":\"https:\/\/yousum.gpucore.co\/#\/schema\/person\/97fa48242daf3908e4d9a5f26f4a059c\",\"name\":\"admin NU\",\"image\":{\"@type\":\"ImageObject\",\"inLanguage\":\"ja\",\"@id\":\"https:\/\/yousum.gpucore.co\/#\/schema\/person\/image\/\",\"url\":\"https:\/\/youzum.net\/wp-content\/uploads\/avatars\/2\/1746849356-bpfull.png\",\"contentUrl\":\"https:\/\/youzum.net\/wp-content\/uploads\/avatars\/2\/1746849356-bpfull.png\",\"caption\":\"admin NU\"},\"url\":\"https:\/\/youzum.net\/ja\/members\/adminnu\/\"}]}<\/script>\n<!-- \/ Yoast SEO plugin. -->","yoast_head_json":{"title":"Hierarchical NeRF with JAX3D for Volumetric Rendering, Novel-View Synthesis, and 3D Reconstruction - YouZum","description":"\u0e01\u0e34\u0e08\u0e01\u0e23\u0e23\u0e21\u0e40\u0e01\u0e35\u0e48\u0e22\u0e27\u0e01\u0e31\u0e1a\u0e42\u0e14\u0e23\u0e19","robots":{"index":"index","follow":"follow","max-snippet":"max-snippet:-1","max-image-preview":"max-image-preview:large","max-video-preview":"max-video-preview:-1"},"canonical":"https:\/\/youzum.net\/ja\/hierarchical-nerf-with-jax3d-for-volumetric-rendering-novel-view-synthesis-and-3d-reconstruction\/","og_locale":"ja_JP","og_type":"article","og_title":"Hierarchical NeRF with JAX3D for Volumetric Rendering, Novel-View Synthesis, and 3D Reconstruction - YouZum","og_description":"\u0e01\u0e34\u0e08\u0e01\u0e23\u0e23\u0e21\u0e40\u0e01\u0e35\u0e48\u0e22\u0e27\u0e01\u0e31\u0e1a\u0e42\u0e14\u0e23\u0e19","og_url":"https:\/\/youzum.net\/ja\/hierarchical-nerf-with-jax3d-for-volumetric-rendering-novel-view-synthesis-and-3d-reconstruction\/","og_site_name":"YouZum","article_publisher":"https:\/\/www.facebook.com\/DroneAssociationTH\/","article_published_time":"2026-09-14T01:45:47+00:00","author":"admin NU","twitter_card":"summary_large_image","twitter_misc":{"\u57f7\u7b46\u8005":"admin NU","\u63a8\u5b9a\u8aad\u307f\u53d6\u308a\u6642\u9593":"18\u5206"},"schema":{"@context":"https:\/\/schema.org","@graph":[{"@type":"Article","@id":"https:\/\/youzum.net\/hierarchical-nerf-with-jax3d-for-volumetric-rendering-novel-view-synthesis-and-3d-reconstruction\/#article","isPartOf":{"@id":"https:\/\/youzum.net\/hierarchical-nerf-with-jax3d-for-volumetric-rendering-novel-view-synthesis-and-3d-reconstruction\/"},"author":{"name":"admin NU","@id":"https:\/\/yousum.gpucore.co\/#\/schema\/person\/97fa48242daf3908e4d9a5f26f4a059c"},"headline":"Hierarchical NeRF with JAX3D for Volumetric Rendering, Novel-View Synthesis, and 3D Reconstruction","datePublished":"2026-09-14T01:45:47+00:00","mainEntityOfPage":{"@id":"https:\/\/youzum.net\/hierarchical-nerf-with-jax3d-for-volumetric-rendering-novel-view-synthesis-and-3d-reconstruction\/"},"wordCount":764,"commentCount":0,"publisher":{"@id":"https:\/\/yousum.gpucore.co\/#organization"},"articleSection":["AI","Committee","News","Uncategorized"],"inLanguage":"ja","potentialAction":[{"@type":"CommentAction","name":"Comment","target":["https:\/\/youzum.net\/hierarchical-nerf-with-jax3d-for-volumetric-rendering-novel-view-synthesis-and-3d-reconstruction\/#respond"]}]},{"@type":"WebPage","@id":"https:\/\/youzum.net\/hierarchical-nerf-with-jax3d-for-volumetric-rendering-novel-view-synthesis-and-3d-reconstruction\/","url":"https:\/\/youzum.net\/hierarchical-nerf-with-jax3d-for-volumetric-rendering-novel-view-synthesis-and-3d-reconstruction\/","name":"Hierarchical NeRF with JAX3D for Volumetric Rendering, Novel-View Synthesis, and 3D Reconstruction - YouZum","isPartOf":{"@id":"https:\/\/yousum.gpucore.co\/#website"},"datePublished":"2026-09-14T01:45:47+00:00","description":"\u0e01\u0e34\u0e08\u0e01\u0e23\u0e23\u0e21\u0e40\u0e01\u0e35\u0e48\u0e22\u0e27\u0e01\u0e31\u0e1a\u0e42\u0e14\u0e23\u0e19","breadcrumb":{"@id":"https:\/\/youzum.net\/hierarchical-nerf-with-jax3d-for-volumetric-rendering-novel-view-synthesis-and-3d-reconstruction\/#breadcrumb"},"inLanguage":"ja","potentialAction":[{"@type":"ReadAction","target":["https:\/\/youzum.net\/hierarchical-nerf-with-jax3d-for-volumetric-rendering-novel-view-synthesis-and-3d-reconstruction\/"]}]},{"@type":"BreadcrumbList","@id":"https:\/\/youzum.net\/hierarchical-nerf-with-jax3d-for-volumetric-rendering-novel-view-synthesis-and-3d-reconstruction\/#breadcrumb","itemListElement":[{"@type":"ListItem","position":1,"name":"Home","item":"https:\/\/youzum.net\/"},{"@type":"ListItem","position":2,"name":"Hierarchical NeRF with JAX3D for Volumetric Rendering, Novel-View Synthesis, and 3D Reconstruction"}]},{"@type":"WebSite","@id":"https:\/\/yousum.gpucore.co\/#website","url":"https:\/\/yousum.gpucore.co\/","name":"YouSum","description":"","publisher":{"@id":"https:\/\/yousum.gpucore.co\/#organization"},"potentialAction":[{"@type":"SearchAction","target":{"@type":"EntryPoint","urlTemplate":"https:\/\/yousum.gpucore.co\/?s={search_term_string}"},"query-input":{"@type":"PropertyValueSpecification","valueRequired":true,"valueName":"search_term_string"}}],"inLanguage":"ja"},{"@type":"Organization","@id":"https:\/\/yousum.gpucore.co\/#organization","name":"Drone Association Thailand","url":"https:\/\/yousum.gpucore.co\/","logo":{"@type":"ImageObject","inLanguage":"ja","@id":"https:\/\/yousum.gpucore.co\/#\/schema\/logo\/image\/","url":"https:\/\/youzum.net\/wp-content\/uploads\/2024\/11\/tranparent-logo.png","contentUrl":"https:\/\/youzum.net\/wp-content\/uploads\/2024\/11\/tranparent-logo.png","width":300,"height":300,"caption":"Drone Association Thailand"},"image":{"@id":"https:\/\/yousum.gpucore.co\/#\/schema\/logo\/image\/"},"sameAs":["https:\/\/www.facebook.com\/DroneAssociationTH\/"]},{"@type":"Person","@id":"https:\/\/yousum.gpucore.co\/#\/schema\/person\/97fa48242daf3908e4d9a5f26f4a059c","name":"admin NU","image":{"@type":"ImageObject","inLanguage":"ja","@id":"https:\/\/yousum.gpucore.co\/#\/schema\/person\/image\/","url":"https:\/\/youzum.net\/wp-content\/uploads\/avatars\/2\/1746849356-bpfull.png","contentUrl":"https:\/\/youzum.net\/wp-content\/uploads\/avatars\/2\/1746849356-bpfull.png","caption":"admin NU"},"url":"https:\/\/youzum.net\/ja\/members\/adminnu\/"}]}},"rttpg_featured_image_url":null,"rttpg_author":{"display_name":"admin NU","author_link":"https:\/\/youzum.net\/ja\/members\/adminnu\/"},"rttpg_comment":0,"rttpg_category":"<a href=\"https:\/\/youzum.net\/ja\/category\/ai-club\/\" rel=\"category tag\">AI<\/a> <a href=\"https:\/\/youzum.net\/ja\/category\/committee\/\" rel=\"category tag\">Committee<\/a> <a href=\"https:\/\/youzum.net\/ja\/category\/news\/\" rel=\"category tag\">News<\/a> <a href=\"https:\/\/youzum.net\/ja\/category\/uncategorized\/\" rel=\"category tag\">Uncategorized<\/a>","rttpg_excerpt":"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&hellip;","_links":{"self":[{"href":"https:\/\/youzum.net\/ja\/wp-json\/wp\/v2\/posts\/117670","targetHints":{"allow":["GET"]}}],"collection":[{"href":"https:\/\/youzum.net\/ja\/wp-json\/wp\/v2\/posts"}],"about":[{"href":"https:\/\/youzum.net\/ja\/wp-json\/wp\/v2\/types\/post"}],"author":[{"embeddable":true,"href":"https:\/\/youzum.net\/ja\/wp-json\/wp\/v2\/users\/2"}],"replies":[{"embeddable":true,"href":"https:\/\/youzum.net\/ja\/wp-json\/wp\/v2\/comments?post=117670"}],"version-history":[{"count":0,"href":"https:\/\/youzum.net\/ja\/wp-json\/wp\/v2\/posts\/117670\/revisions"}],"wp:attachment":[{"href":"https:\/\/youzum.net\/ja\/wp-json\/wp\/v2\/media?parent=117670"}],"wp:term":[{"taxonomy":"category","embeddable":true,"href":"https:\/\/youzum.net\/ja\/wp-json\/wp\/v2\/categories?post=117670"},{"taxonomy":"post_tag","embeddable":true,"href":"https:\/\/youzum.net\/ja\/wp-json\/wp\/v2\/tags?post=117670"}],"curies":[{"name":"wp","href":"https:\/\/api.w.org\/{rel}","templated":true}]}}