AI サービスが一時的に利用できないため、復旧後に翻訳を補完します。
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]