nvalchemi-dynamics-hooks

작성자: nvidia

다이나믹스 훅을 사용하고 작성하는 방법 — 각 시뮬레이션 단계 중 특정 지점에서 배치 상태를 관찰하거나 수정하는 콜백입니다. 시뮬레이션이 필요할 때 사용합니다…

npx skills add https://github.com/nvidia/nvalchemi-toolkit --skill nvalchemi-dynamics-hooks

nvalchemi Hooks

Overview

Hooks are callbacks that fire at specific points during each workflow step. They observe or modify batch state without changing the engine itself. The hook system is framework-wide: the same Hook protocol works for dynamics and custom pipelines. Dynamics engines pass DynamicsContext; custom engines can pass HookContext or their own context subclass.

from nvalchemi.hooks import (
    BiasedPotentialHook,
    DynamicsContext,
    Hook,
    HookContext,
    HookRegistryMixin,
    NeighborListHook,
    WrapPeriodicHook,
)
from nvalchemi.dynamics.base import DynamicsStage
from nvalchemi.dynamics.hooks import (
    EnergyDriftMonitorHook,
    LoggingHook,
    MaxForceClampHook,
    NaNDetectorHook,
    SnapshotHook,
    StageTimingHook,
    TorchProfilerHook,
)

Hook protocol

Any object with these attributes satisfies the Hook protocol (runtime-checkable):

class Hook(Protocol):
    frequency: int        # execute every N steps (1 = every step)
    stage: Enum | None    # stage enum value (None for stage-agnostic hooks)

    def __call__(self, ctx: HookContext, stage: Enum) -> None:
        """Called with a context snapshot and the current stage."""
        ...

A hook fires when step_count % hook.frequency == 0 (so all hooks fire at step 0).

HookContext — base snapshot shared by hook-enabled workflows:

@dataclass(kw_only=True)
class HookContext:
    batch: Batch              # current batch (all engines)
    model: BaseModelMixin | None = None
    global_rank: int = 0      # distributed rank
    workflow: Any = None      # back-reference to the engine

DynamicsContext — context passed by dynamics engines:

@dataclass(kw_only=True)
class DynamicsContext(HookContext):
    step_count: int = 0
    converged_mask: torch.Tensor | None = None

Access batch data via ctx.batch and dynamics step info via ctx.step_count.


Execution stages

Dynamics — DynamicsStage

Each step() call fires hooks at 9 stages in this order:

BEFORE_STEP (0)
  BEFORE_PRE_UPDATE (1)  →  pre_update()  →  AFTER_PRE_UPDATE (2)
  BEFORE_COMPUTE (3)     →  compute()      →  AFTER_COMPUTE (4)
  BEFORE_POST_UPDATE (5) →  post_update()  →  AFTER_POST_UPDATE (6)
AFTER_STEP (7)
ON_CONVERGE (8)   ← only if convergence detected

Stage selection guidelines (dynamics):

GoalStage
Modify forces/energy after modelDynamicsStage.AFTER_COMPUTE
Observe final state (logging, snapshots)DynamicsStage.AFTER_STEP
Wrap positions after velocity updateDynamicsStage.AFTER_POST_UPDATE
Instrument timing / profilingDynamicsStage.BEFORE_STEP
React to convergenceDynamicsStage.ON_CONVERGE

Registering hooks

from nvalchemi.dynamics.demo import DemoDynamics

# At construction
dynamics = DemoDynamics(
    model=model,
    n_steps=1000,
    dt=0.5,
    hooks=[
        MaxForceClampHook(max_force=10.0),
        LoggingHook(backend="csv", log_path="md_log.csv", frequency=100),
    ],
)

# After construction
dynamics.register_hook(NaNDetectorHook(frequency=10))

Multiple hooks at the same stage fire in registration order.

Stage type enforcement: each engine declares _stage_type to restrict which enum types are accepted. For example, BaseDynamics sets _stage_type = DynamicsStage.


Built-in hooks

Safety hooks (stage: AFTER_COMPUTE)

NaNDetectorHook — detect NaN/Inf in forces and energy.

NaNDetectorHook(
    frequency=1,              # check every N steps
    extra_keys=["stress"],    # additional batch keys to check (optional)
)

MaxForceClampHook — clamp per-atom force vectors to a maximum L2 norm.

MaxForceClampHook(
    max_force=10.0,     # max force norm (eV/A)
    frequency=1,
)

Bias hook (stage: AFTER_COMPUTE)

BiasedPotentialHook — add an external bias potential for enhanced sampling.

def my_bias(batch: Batch) -> tuple[torch.Tensor, torch.Tensor]:
    """Return (bias_energy [B, 1], bias_forces [V, 3])."""
    bias_e = torch.zeros(batch.num_graphs, 1, device=batch.device)
    bias_f = torch.zeros_like(batch.positions)
    # ... compute bias ...
    return bias_e, bias_f

BiasedPotentialHook(
    bias_fn=my_bias,
    stage=DynamicsStage.AFTER_COMPUTE,
    frequency=1,
)

Observer hooks (stage: AFTER_STEP)

LoggingHook — log scalar observables.

LoggingHook(
    backend="csv",                  # "csv", "tensorboard", or "custom"
    frequency=100,
    log_path="md_log.csv",          # for file-based backends
    custom_scalars={                # additional scalars to log
        "max_velocity": lambda ctx: ctx.batch.velocities.norm(dim=-1).max(),
    },
    writer_fn=None,                 # custom writer for "custom" backend
)

SnapshotHook — save full batch state to a DataSink.

from nvalchemi.dynamics.sinks import GPUBuffer, HostMemory, ZarrData

SnapshotHook(
    sink=ZarrData("trajectory.zarr", capacity=10000),
    frequency=10,
)

EnergyDriftMonitorHook — track total energy drift.

EnergyDriftMonitorHook(
    threshold=1e-4,                          # drift threshold
    metric="per_atom_per_step",              # or "absolute"
    action="warn",                           # or "raise"
    frequency=1,
    include_kinetic=True,                    # include kinetic energy
)

Periodic boundary hook (stage: AFTER_POST_UPDATE)

WrapPeriodicHook — wrap positions back into the unit cell.

WrapPeriodicHook(frequency=10, stage=DynamicsStage.AFTER_POST_UPDATE)

Profiling hooks (multi-stage)

StageTimingHook — per-stage NVTX ranges and wall-clock timing. Registers itself at every profiled stage via _runs_on_stage, records timestamps, and computes per-transition deltas (optionally written to CSV or console).

StageTimingHook(
    profiled_stages="all",                  # "all", "step", "detailed", or a set[Enum]
    frequency=1,
    enable_nvtx=True,                       # NVTX push/pop ranges for Nsight Systems
    timer_backend="auto",                   # "auto", "cuda_event", or "perf_counter"
    log_path="timing.csv",                  # optional CSV of per-transition timings
    show_console=False,                     # print a timing table via loguru
)

Call profiler.summary() after the run for aggregated per-stage timings. For full kernel-level PyTorch profiler traces, use TorchProfilerHook, which captures traces through PhysicsNeMo's profiler wrapper.


Writing a custom hook

Option 1: Simple single-stage hook (dynamics)

Implement the protocol directly — no inheritance needed.

from nvalchemi.dynamics.base import DynamicsStage
from nvalchemi.hooks import DynamicsContext

class TemperatureLogger:
    stage = DynamicsStage.AFTER_STEP
    frequency = 50

    def __call__(self, ctx: DynamicsContext, stage: DynamicsStage) -> None:
        ke = ctx.batch.kinetic_energies.sum()
        n_atoms = ctx.batch.num_nodes
        temp = 2.0 * ke / (3.0 * n_atoms * 8.617e-5)  # kB in eV/K
        print(f"Step {ctx.step_count}: T = {temp:.1f} K")

Option 2: Multi-stage hook with _runs_on_stage

Fire at multiple stages by defining _runs_on_stage(stage) -> bool:

from enum import Enum
from nvalchemi.dynamics.base import DynamicsStage
from nvalchemi.hooks import DynamicsContext

class StepTimerHook:
    stage = DynamicsStage.BEFORE_STEP  # primary stage (protocol compliance)
    frequency = 1

    def __init__(self):
        self._stages = {DynamicsStage.BEFORE_STEP, DynamicsStage.AFTER_STEP}
        self._t0 = None

    def _runs_on_stage(self, stage: Enum) -> bool:
        return stage in self._stages

    def __call__(self, ctx: DynamicsContext, stage: DynamicsStage) -> None:
        import time
        if stage == DynamicsStage.BEFORE_STEP:
            self._t0 = time.perf_counter()
        elif stage == DynamicsStage.AFTER_STEP and self._t0 is not None:
            dt = time.perf_counter() - self._t0
            print(f"Step {ctx.step_count}: {dt*1000:.1f} ms")

Option 3: Cross-category hook with plum dispatch

For hooks that work with multiple stage enum types (e.g. DynamicsStage and a custom enum), use plum.dispatch to overload __call__ with different stage types:

from dataclasses import dataclass
from enum import Enum
from plum import dispatch
from nvalchemi.dynamics.base import DynamicsStage
from nvalchemi.hooks import DynamicsContext, HookContext

# Example custom stage enum for a hypothetical pipeline
class MyPipelineStage(Enum):
    BEFORE_PROCESS = 0
    AFTER_PROCESS = 1


@dataclass(kw_only=True)
class PipelineContext(HookContext):
    step_count: int = 0


class UniversalLoggerHook:
    stage = DynamicsStage.AFTER_STEP
    frequency = 10

    def __init__(self):
        self._stages = {DynamicsStage.AFTER_STEP, MyPipelineStage.AFTER_PROCESS}

    def _runs_on_stage(self, stage: Enum) -> bool:
        return stage in self._stages

    @dispatch
    def __call__(self, ctx: DynamicsContext, stage: DynamicsStage) -> None:
        fmax = ctx.batch.forces.norm(dim=-1).max().item()
        print(f"[dynamics] step {ctx.step_count}: fmax={fmax:.4f}")

    @dispatch
    def __call__(self, ctx: PipelineContext, stage: MyPipelineStage) -> None:
        print(f"[pipeline] step {ctx.step_count}: processed")

    @dispatch
    def __call__(self, ctx: HookContext, stage: Enum) -> None:
        print(f"[custom] stage={stage.name}, graphs={ctx.batch.num_graphs}")

Use this plum.dispatch pattern when one hook must handle several context/stage types. Built-in multi-stage hooks like StageTimingHook instead use the simpler _runs_on_stage approach from Option 2.


Hook ordering recommendations

Register hooks in this order for correct behavior:

hooks = [
    # 1. Bias (modifies forces/energy)
    BiasedPotentialHook(bias_fn=my_bias, stage=DynamicsStage.AFTER_COMPUTE),
    # 2. Safety (clamp after all force modifications)
    MaxForceClampHook(max_force=10.0),
    # 3. NaN detection (check final forces)
    NaNDetectorHook(),
    # 4. Periodic wrapping
    WrapPeriodicHook(frequency=10, stage=DynamicsStage.AFTER_POST_UPDATE),
    # 5. Observers (read final state)
    LoggingHook(backend="csv", log_path="md_log.csv", frequency=100),
    SnapshotHook(sink=my_sink, frequency=50),
    EnergyDriftMonitorHook(threshold=1e-4),
    # 6. Profiling
    StageTimingHook(),
]

dynamics = DemoDynamics(model=model, n_steps=10000, dt=0.5, hooks=hooks)

Complete example

import torch
from nvalchemi.data import AtomicData, Batch
from nvalchemi.models.demo import DemoModel, DemoModelWrapper
from nvalchemi.dynamics.demo import DemoDynamics
from nvalchemi.dynamics.base import DynamicsStage
from nvalchemi.hooks import DynamicsContext
from nvalchemi.dynamics.hooks import MaxForceClampHook, NaNDetectorHook

# Custom hook
class StepPrinter:
    stage = DynamicsStage.AFTER_STEP
    frequency = 10

    def __call__(self, ctx: DynamicsContext, stage: DynamicsStage) -> None:
        fmax = ctx.batch.forces.norm(dim=-1).max().item()
        print(f"Step {ctx.step_count}: fmax={fmax:.4f}")

# Setup
model = DemoModelWrapper(DemoModel())
dynamics = DemoDynamics(
    model=model,
    n_steps=100,
    dt=0.5,
    hooks=[
        MaxForceClampHook(max_force=10.0),
        NaNDetectorHook(),
        StepPrinter(),
    ],
)

data = AtomicData(
    atomic_numbers=torch.tensor([6, 6, 8], dtype=torch.long),
    positions=torch.randn(3, 3),
)
batch = Batch.from_data_list([data])
batch.forces = torch.zeros(3, 3)
batch.energy = torch.zeros(1, 1)

dynamics.run(batch)

nvidia의 다른 스킬

compileiq-debug
nvidia
무언가 잘못되었을 때 사용: Search()가 멈추거나, 모든 평가가 INVALID_SCORE를 반환하거나, 점수가 개선되지 않거나, 모든 설정이 동일한 숫자를 반환하거나, ptxas 오류 등이 발생할 때
create-github-pr
nvidia
gh CLI를 사용하여 GitHub 풀 리퀘스트를 생성합니다. 사용자가 새 PR을 만들거나, 코드 리뷰를 제출하거나, 풀 리퀘스트를 열고자 할 때 사용합니다. 트리거 키워드 -…
nemoclaw-maintainer-cross-issue-sweep
nvidia
다른 열린 이슈들을 스캔하여 주어진 PR이 함께 수정하거나 실수로 망가뜨릴 수 있는 이슈를 찾습니다. 인접 수정 기회와 모순 위험을 file:line…과 함께 출력합니다.
fhir-basics
nvidia
에이전트에게 FHIR R4 API의 작동 방식, 사용 가능한 리소스, 검색 매개변수를 사용한 쿼리 방법, 모든 응답 형식을 올바르게 파싱하는 방법을 가르칩니다…
compileiq-validate-result
nvidia
검색이 완료된 후, 속도 향상을 청구하거나 ACF를 발송하기 전에 사용합니다. dump_results CSV를 로드하고, 상위 K개 후보(단일 목표)를 추출합니다…
changelog-audit
nvidia
릴리스 전에 Warp CHANGELOG.md를 감사합니다: 누락된 항목 복구, 사용자 영향별 정렬, 항목 언어 다듬기, 줄 바꿈, (릴리스 브랜치 모드) 비교 업데이트…
maintain-dynamic-plugins
nvidia
NeMo Relay 동적 플러그인 로더, 매니페스트, Rust 네이티브 SDK, gRPC 워커 프로토콜, Python 워커 SDK, 문서, 테스트 및 릴리스 워크플로 커버리지를 유지 관리합니다.
dgx-diagnose
nvidia
일반적인 DGX Station GB300 문제 진단 — CUDA 충돌, 잘못된 GPU 타겟팅, vLLM/SGLang 컨테이너 버그, MIG 상태 문제, NVLink/Fabric Manager 오류,…