В мире компьютерного зрения и генеративного искусственного интеллекта Neural Radiance Fields (NeRF) стали золотым стандартом для синтеза новых ракурсов и 3D-реконструкции. Однако классические реализации NeRF часто страдают от вычислительной сложности и медленной сходимости. В этой статье мы подробно разберем продвинутый подход: построение иерархического NeRF с использованием мощного стека JAX, Flax, Optax и специализированных примитивов объемного рендеринга из библиотеки jax3d. Это не просто набор инструкций, а глубокое погружение в то, как современные фреймворки позволяют создавать высококачественные 3D-модели из 2D-изображений, используя параллельные вычисления и эффективное сэмплирование.
Наш путь будет полным циклом: от создания аналитической сцены с объемной геометрией и зависимым от ракурса излучением до финальной оценки качества через метрики PSNR, визуализацию глубины и извлечения полигональной сетки. Мы рассмотрим, как работают функции сэмплирования лучей, как устроена архитектура нейросети с пропускными связями (skip connections) и почему иерархический подход (coarse-to-fine) критически важен для качества результата. Все примеры кода адаптированы для работы в среде Jupyter/Colab, с учетом особенностей запуска как на GPU, так и на CPU.
011. Подготовка среды и загрузка библиотек JAX3D
Первый шаг в любом проекте на базе JAX — это корректная настройка окружения. Библиотека jax3d от Google Research предоставляет оптимизированные примитивы для работы с 3D-данными, включая функции для объемного рендеринга, которые являются ядром нашего проекта. Важно отметить, что jax3d не всегда устанавливается как отдельный пакет через pip в стандартном виде, поэтому мы используем подход прямого импорта модулей из репозитория.
Для начала необходимо установить зависимости. Мы используем etils для работы с путями и типами данных, flax для построения нейронных сетей, optax для оптимизации и scikit-image для обработки изображений. Код ниже демонстрирует установку пакетов и клонирование репозитория jax3d, если он еще не присутствует в системе.
import os
import sys
import subprocess
import importlib.util
import functools
import dataclasses
import time
import 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]" 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}\nTry: 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(f"jax {jax.__version__} | device: {jax.devices()[0].device_kind} ({jax.devices()[0].platform})")
print("jax3d volume_rendering API:")
for func in ["sample_along_rays", "volume_rendering", "sample_piecewise_constant_pdf", "sample_1d"]:
if hasattr(j3vr, func):
print(f" - {func}")022. Конфигурация и генерация синтетической сцены
Прежде чем обучать нейросеть, нам нужны данные. В данном туториале мы не используем реальные фотографии, а создаем аналитическую сцену — набор математически заданных объектов (сферы, пол), которые имеют четкую геометрию и физические свойства (албедо, шероховатость). Это позволяет нам иметь "ground truth" (истинное значение) для сравнения.

Ключевым элементом здесь является функция orbit_poses, которая генерирует позиции камер. Мы используем "golden-angle" распределение азимута и монотонные изменения возвышения, чтобы камеры равномерно покрывали сцену, подобно куполу. Это обеспечивает хорошее покрытие углов обзора, необходимое для качественной реконструкции.
@dataclasses.dataclass
class Config:
H: int = 64
W: int = 64
n_train_views: int = 24
n_test_views: int = 24
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 = 2
width: int = 128
depth: int = 4
skip: int = 4
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=2, skip=2, batch_rays=1024, chunk=1600, grid_res=64)Генерация лучей (rays_from_pose) использует пинхольную модель камеры. Мы вычисляем направления лучей для каждого пикселя изображения, преобразуем их в мировое пространство с помощью матрицы камеры-to-world (c2w) и получаем начальные точки (origins). Эти лучи являются основой для последующего сэмплирования внутри сцены.
_sphere_field для сфер с бликами (specular) и _floor_field для пола с шахматным паттерном. Это позволяет модели учить не только геометрию, но и физические свойства материалов, такие как отражение света.033. Архитектура NeRF: От позиции к цвету
Сердце NeRF — это нейронная сеть, которая принимает на вход 3D-координату точки и направление взгляда, а на выходе выдает объемную плотность (sigma) и цвет (rgb). В нашей реализации используется архитектура на базе Flax (JAX-версия Flax Linen).
Ключевые особенности архитектуры:

- Синусоидальное позиционное кодирование (PosEnc): Входные координаты (позиция и направление) преобразуются с помощью синусоидальных функций разных частот. Это помогает нейросети улавливать высокочастотные детали сцены, которые трудно аппроксимировать обычными слоями.
- Разделение геометрии и внешнего вида: Сеть имеет два выхода. Первый выход (через
softplus) гарантирует неотрицательную плотность, отвечающую за форму объектов. Второй выход зависит от направления взгляда, что позволяет моделировать зеркальные отражения и другие эффекты, зависящие от ракурса. - Skip Connections: На каждом
skip-м слое входные закодированные данные добавляются к выходу слоя. Это улучшает поток градиентов и помогает сети сохранять информацию о исходных координатах.
def posenc(deg):
"""NeRF sinusoidal encoding, with the raw input concatenated."""
if deg == 0:
return
scales = 2.0 ** (jnp.arange(deg, dtype=dtype) * 2 * jnp.pi)
xb = jnp.broadcast_to(x[..., None], shape + (deg,)) * scales
return jnp.concatenate([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)
for i in range(self.depth):
h = nn.relu(nn.Dense(self.width)(inp))
if i % self.skip == 0:
inp = jnp.concatenate([inp, h], axis=-1)
sigma = nn.softplus(nn.Dense(1)(h))
h = nn.Dense(self.width)(h)
h = nn.relu(h)
h = jnp.concatenate([h, posenc(dirs, self.deg_dir)], axis=-1)
rgb = nn.sigmoid(nn.Dense(3)(h))
return sigma, rgb
model = NeRFMLP(
width=cfg.width,
depth=cfg.depth,
skip=cfg.skip,
deg_pos=cfg.deg_pos,
deg_dir=cfg.deg_dir
)044. Иерархический рендеринг: Coarse и Fine проходы
Простое сэмплирование точек вдоль луча часто приводит к артефактам или требует огромного количества точек для качества. Решение — иерархический подход. Мы выполняем два прохода:
- Coarse Pass (Грубый проход): Мы равномерно сэмплируем
n_coarseточек вдоль луча с помощьюj3vr.sample_along_rays. Затем мы используемj3vr.volume_renderingдля композитинга этих точек. Этот проход дает нам предварительное представление о сцене и, что важнее, веса (importance weights) для каждой точки. - Importance Sampling: Используя веса из грубого прохода, мы строим кусочно-постоянное распределение вероятностей (PDF). Затем мы сэмплируем новые, более密集ные точки в тех областях луча, где нейросеть предсказала высокую плотность или цвет. Это делается через
j3vr.sample_piecewise_constant_pdf. - Fine Pass (Тонкий проход): Мы объединяем старые и новые точки, сортируем их и подаем в ту же сеть (или вторую сеть) для финального рендеринга. Градиенты через операцию сэмплирования отключаются (
jax.lax.stop_gradient), чтобы обучение было стабильным.
def render_rays(params, origins, dirs, rng, deterministic=False):
"""Coarse pass -> importance-resample -> fine pass."""
rng_c, rng_f = jax.random.split(rng)
# Coarse pass
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, pos_c.shape + (3,))
sigma_c, rgb_c = model.apply(params, "coarse", pos_c, dirs_c)
out_c = j3vr.volume_rendering(
sample_values="rgb", rgb=rgb_c, sample_density=sigma_c,
depths=depths_c, background_values="rgb", background_value=jnp.ones_like(rgb_c)
)
# Importance sampling
mid = 0.5 * (depths_c[..., 1:] + depths_c[..., :-1])
bin_edges = jnp.concatenate([depths_c[..., :1], mid, depths_c[..., -1:]], axis=-1)
t_fine = j3vr.sample_piecewise_constant_pdf(
bin_edges, bin_edges, weights=out_c.sample_weights,
sample_count=cfg.n_fine, deterministic=deterministic, rng=rng_f
)
t_fine = jax.lax.stop_gradient(t_fine)
# Fine pass
depths_f = jnp.sort(jnp.concatenate([depths_c, t_fine], axis=-1), axis=-1)
pos_f = origins[..., None, :] + depths_f[..., None] * dirs[..., None, :]
dirs_f = jnp.broadcast_to(dirs, pos_f.shape + (3,))
sigma_f, rgb_f = model.apply(params, "fine", pos_f, dirs_f)
out_f = j3vr.volume_rendering(
sample_values="rgb", rgb=rgb_f, sample_density=sigma_f,
depths=depths_f, background_values="rgb", background_value=jnp.ones_like(rgb_f)
)
aux = {"depths_c": depths_c, "weights_c": out_c.sample_weights, "t_fine": t_fine}
return out_c, out_f, aux055. Обучение модели: Оптимизация и JIT
Обучение NeRF — это процесс минимизации разницы между предсказанным цветом и истинным цветом пикселя (MSE loss). Мы используем оптимизатор Adam с экспоненциальным затуханием скорости обучения (learning rate decay). Это позволяет модели быстро сходиться в начале и тонко настраивать детали в конце.

Ключевая оптимизация производительности — использование @jax.jit. JAX компилирует функцию train_step в эффективный XLA-код, что значительно ускоряет вычисления, особенно на GPU. Мы также применяем градиентное клиппирование (clip_by_global_norm), чтобы предотвратить взрыв градиентов, характерный для глубоких сетей.
schedule = optax.exponential_decay(cfg.lr_init, cfg.steps, cfg.lr_final / cfg.lr_init)
tx = optax.chain(optax.clip_by_global_norm(1.0), optax.adam(schedule))
state = train_state.TrainState.create(apply_fn=model.apply, params=params, tx=tx)
@jax.jit
def train_step(state, target, rng):
def loss_fn(params):
out_c, out_f, _ = render_rays(params, rng, deterministic=False)
l_c = jnp.mean((out_c.ray_values["rgb"] - target) ** 2)
l_f = jnp.mean((out_f.ray_values["rgb"] - target) ** 2)
return l_c + l_f, l_f
loss, l_fine = jax.value_and_grad(loss_fn, has_aux=True)(state.params)
state = state.apply_gradients(grads=loss.grads)
return state, loss, l_fineВо время обучения мы мониторим метрику PSNR (Peak Signal-to-Noise Ratio). График PSNR от шага обучения показывает, как модель улучшает качество реконструкции. Хороший результат для синтетических сцен часто превышает 30-35 dB.
066. Оценка и визуализация результатов
После обучения мы оцениваем модель на тестовых видах, которые не использовались в процессе обучения. Мы рендерим изображения, карты глубины и альфа-каналы (непрозрачность). Это позволяет не только оценить визуальное качество, но и понять, как модель воспринимает структуру сцены.
Для оценки качества мы используем PSNR. Также важно визуализировать глубину: в NeRF глубина вычисляется как средневзвешенное расстояние точек сэмплирования, где веса — это непрозрачность. Это дает более точную оценку геометрии, чем просто расстояние до первой непрозрачной точки.

@jax.jit
def render_chunk(params, rng, origins, dirs):
out_f, aux = render_rays(params, origins, dirs, rng, deterministic=True)
return out_f.ray_values["rgb"], out_f.ray_depth, out_f.ray_alpha, aux
def render_image(params, origins, dirs, rng):
origins = jnp.asarray(origins).reshape(-1, 3)
dirs = jnp.asarray(dirs).reshape(-1, 3)
rgb_list, dep_list, alp_list = [], [], []
for i in range(0, origins.shape[0], cfg.chunk):
oc, dc, ac, _ = render_chunk(params, rng, origins[i:i+cfg.chunk], dirs[i:i+cfg.chunk])
rgb_list.append(oc)
dep_list.append(dc)
alp_list.append(ac)
return np.asarray(jnp.concatenate(rgb_list).reshape(cfg.H, cfg.W, 3)), \
np.asarray(jnp.concatenate(dep_list).reshape(cfg.H, cfg.W)), \
np.asarray(jnp.concatenate(alp_list).reshape(cfg.H, cfg.W))077. Извлечение 3D-геометрии: Marching Cubes
Одним из самых мощных преимуществ NeRF является возможность извлечения 3D-мешей (mesh). Поскольку нейросеть предсказывает плотность (sigma), мы можем рассматривать сцену как объемную функцию. Используя алгоритм Marching Cubes, мы можем найти изоповерхность, где плотность равна определенному порогу, и получить полигональную сетку.
Для этого мы создаем регулярную сетку в пространстве сцены, подаем координаты точек в обученную модель NeRF, получаем значения плотности и применяем алгоритм извлечения поверхности. Это позволяет получить интерактивную 3D-модель, которую можно экспортировать в форматы OBJ или PLY для использования в других приложениях.
08Что это значит на практике
Реализация иерархического NeRF на JAX3D демонстрирует переход от теоретических концепций к практическим инструментам компьютерного зрения. Использование JAX обеспечивает скорость и масштабируемость, необходимые для обучения сложных моделей. Иерархическое сэмплирование решает проблему качества реконструкции, позволяя получать фотореалистичные изображения с новых ракурсов. А возможность извлечения 3D-геометрии открывает двери для интеграции таких моделей в VR/AR приложения, робототехнику и цифровые двойники.
Для разработчиков в России и СНГ этот подход особенно актуален, так как позволяет запускать вычисления локально, без зависимости от облачных сервисов, используя открытые библиотеки. Освоение этих техник дает глубокое понимание того, как современные AI-системы "видят" и "понимают" трехмерный мир.
Источник: MarkTechPost ↗
