Главная/Блог/Гайд/Иерархический NeRF на JAX3D: Полный…
Гайд11 мин чтения · 14 сентября 2026 г.

Иерархический NeRF на JAX3D: Полный гайд по 3D-реконструкции

Разбираем архитектуру иерархического Neural Radiance Field на базе JAX3D, Flax и Optax. От генерации синтетических данных до извлечения 3D-геометрии через Marching Cubes.

Иерархический NeRF на JAX3D: Полный гайд по 3D-реконструкции

В мире компьютерного зрения и генеративного искусственного интеллекта 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, если он еще не присутствует в системе.

terminalpython
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}")
💡
Важно для локального запуска. Если вы запускаете код на локальной машине, убедитесь, что у вас установлена последняя версия JAX, совместимая с вашим GPU (CUDA/cuDNN). Для CPU-режима код автоматически переключится на уменьшенные параметры, чтобы избежать переполнения памяти, но обучение будет значительно медленнее.

022. Конфигурация и генерация синтетической сцены

Прежде чем обучать нейросеть, нам нужны данные. В данном туториале мы не используем реальные фотографии, а создаем аналитическую сцену — набор математически заданных объектов (сферы, пол), которые имеют четкую геометрию и физические свойства (албедо, шероховатость). Это позволяет нам иметь "ground truth" (истинное значение) для сравнения.

Иерархический NeRF на JAX3D: Полный гайд по 3D-реконструкции

Ключевым элементом здесь является функция orbit_poses, которая генерирует позиции камер. Мы используем "golden-angle" распределение азимута и монотонные изменения возвышения, чтобы камеры равномерно покрывали сцену, подобно куполу. Это обеспечивает хорошее покрытие углов обзора, необходимое для качественной реконструкции.

terminalpython
@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).

Ключевые особенности архитектуры:

Иерархический NeRF на JAX3D: Полный гайд по 3D-реконструкции
  1. Синусоидальное позиционное кодирование (PosEnc): Входные координаты (позиция и направление) преобразуются с помощью синусоидальных функций разных частот. Это помогает нейросети улавливать высокочастотные детали сцены, которые трудно аппроксимировать обычными слоями.
  2. Разделение геометрии и внешнего вида: Сеть имеет два выхода. Первый выход (через softplus) гарантирует неотрицательную плотность, отвечающую за форму объектов. Второй выход зависит от направления взгляда, что позволяет моделировать зеркальные отражения и другие эффекты, зависящие от ракурса.
  3. Skip Connections: На каждом skip-м слое входные закодированные данные добавляются к выходу слоя. Это улучшает поток градиентов и помогает сети сохранять информацию о исходных координатах.
terminalpython
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 проходы

Простое сэмплирование точек вдоль луча часто приводит к артефактам или требует огромного количества точек для качества. Решение — иерархический подход. Мы выполняем два прохода:

  1. Coarse Pass (Грубый проход): Мы равномерно сэмплируем n_coarse точек вдоль луча с помощью j3vr.sample_along_rays. Затем мы используем j3vr.volume_rendering для композитинга этих точек. Этот проход дает нам предварительное представление о сцене и, что важнее, веса (importance weights) для каждой точки.
  2. Importance Sampling: Используя веса из грубого прохода, мы строим кусочно-постоянное распределение вероятностей (PDF). Затем мы сэмплируем новые, более密集ные точки в тех областях луча, где нейросеть предсказала высокую плотность или цвет. Это делается через j3vr.sample_piecewise_constant_pdf.
  3. Fine Pass (Тонкий проход): Мы объединяем старые и новые точки, сортируем их и подаем в ту же сеть (или вторую сеть) для финального рендеринга. Градиенты через операцию сэмплирования отключаются (jax.lax.stop_gradient), чтобы обучение было стабильным.
terminalpython
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, aux
⚠️
Почему это важно? Без importance sampling нейросеть может "пропустить" тонкие объекты или границы, если они не попали в равномерную сетку точек. Иерархический подход фокусирует вычислительные ресурсы на тех участках луча, где действительно происходит взаимодействие света с объектами.

055. Обучение модели: Оптимизация и JIT

Обучение NeRF — это процесс минимизации разницы между предсказанным цветом и истинным цветом пикселя (MSE loss). Мы используем оптимизатор Adam с экспоненциальным затуханием скорости обучения (learning rate decay). Это позволяет модели быстро сходиться в начале и тонко настраивать детали в конце.

Иерархический NeRF на JAX3D: Полный гайд по 3D-реконструкции

Ключевая оптимизация производительности — использование @jax.jit. JAX компилирует функцию train_step в эффективный XLA-код, что значительно ускоряет вычисления, особенно на GPU. Мы также применяем градиентное клиппирование (clip_by_global_norm), чтобы предотвратить взрыв градиентов, характерный для глубоких сетей.

terminalpython
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 глубина вычисляется как средневзвешенное расстояние точек сэмплирования, где веса — это непрозрачность. Это дает более точную оценку геометрии, чем просто расстояние до первой непрозрачной точки.

Иерархический NeRF на JAX3D: Полный гайд по 3D-реконструкции
terminalpython
@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 ↗