nvalchemi-loss-api

द्वारा nvidia

बिल्ट-इन लॉस फंक्शन का उपयोग कैसे करें और BaseLossFunction टेम्पलेट-मेथड पैटर्न का उपयोग करके कस्टम लॉस कैसे लागू करें — रेसिड्यूअल प्रकार, प्रति-परमाणु सामान्यीकरण,…

npx skills add https://github.com/nvidia/nvalchemi-toolkit --skill nvalchemi-loss-api

nvalchemi Loss API

Overview

Loss functions are torch.nn.Module subclasses rooted at BaseLossFunction. Each leaf consumes (pred, target, **kwargs) and returns a scalar. ComposedLossFunction routes keyed prediction/target mappings to leaves, applies per-component weights (float or LossWeightSchedule), and returns a ComposedLossOutput TypedDict.

from nvalchemi.training import (
    BaseLossFunction,
    ComposedLossFunction,
    ReductionContext,
    EnergyMSELoss,
    EnergyMAELoss,
    ForceMSELoss,
    ForceL2NormLoss,
    StressMSELoss,
)

Built-in losses

Choose losses by the training signal you want:

  • EnergyMSELoss: default for smooth energy regression when larger errors should dominate early training; combine with per_atom=True when system sizes vary.
  • EnergyMAELoss: more robust to outlier energies and often useful for reporting or late-stage fitting when median absolute accuracy matters.
  • EnergyHuberLoss: compromise between MSE and MAE; use when energy labels have occasional noisy outliers but small residuals should remain smooth.
  • ForceMSELoss: default force objective; component-wise squared residuals give strong gradients for geometry-sensitive fitting.
  • ForceL2NormLoss: use when vector direction/magnitude per atom is the desired error signal rather than independent xyz components.
  • ForceHuberLoss: robust force fitting when some force labels are noisy or contain rare large residuals.
  • StressMSELoss / StressHuberLoss: add only when stress labels are reliable and the model is configured to produce stresses.

Composition sugar:

loss_fn = 1.0 * EnergyMSELoss() + 10.0 * ForceMSELoss() + 0.1 * StressMSELoss()
out = loss_fn(predictions, targets, step=step, epoch=epoch, batch=batch)
out["total_loss"].backward()

Graph metadata: losses that need graph structure (per_atom=True, normalize_by_atom_count=True, or padded layouts) accept batch= (pulls batch_idx, num_graphs, num_nodes_per_graph automatically) or explicit kwargs.


Template-method pattern

BaseLossFunction.forward orchestrates five hooks:

forward(pred, target, **kwargs)
  1. validate(pred, target)                         # shape checks
  2. pred, target, ctx = normalize(pred, target, **kwargs)  # pre-processing
  3. valid = mask(pred, target, ctx, **kwargs)       # boolean validity mask
  4. residual = compute_residual(pred, target, valid) # ABSTRACT — must override
  5. scalar = reduce(residual, valid, ctx, **kwargs)  # collapse to scalar

Minimum implementation: override compute_residual only. Defaults handle shape validation, all-True masking, and validity-weighted mean reduction.


Writing a custom loss

Minimal: compute_residual only

class HuberEnergyLoss(BaseLossFunction):
    def __init__(self, *, target_key="energy", prediction_key="predicted_energy", delta=1.0):
        super().__init__()
        self.target_key = target_key
        self.prediction_key = prediction_key
        self.delta = delta

    def compute_residual(self, pred, target, valid):
        residual = torch.where(valid, pred - target, torch.zeros_like(pred))
        abs_r = residual.abs()
        return torch.where(
            abs_r < self.delta,
            0.5 * residual.pow(2),
            self.delta * (abs_r - 0.5 * self.delta),
        )

Per-atom normalization (normalize override)

Override normalize to divide by atom counts and pass weights via ReductionContext["weights"]. The base reduce picks up weights automatically.

class PerAtomEnergyMSE(BaseLossFunction):
    target_key = "energy"
    prediction_key = "predicted_energy"

    def normalize(self, pred, target, **kwargs):
        ctx = ReductionContext()
        counts = kwargs["num_nodes_per_graph"].to(dtype=pred.dtype).unsqueeze(-1).clamp_min(1.0)
        ctx["weights"] = counts  # base reduce uses this for atom-count-weighted mean
        return pred / counts, target / counts, ctx

    def compute_residual(self, pred, target, valid):
        residual = torch.where(valid, pred - target, torch.zeros_like(pred))
        return residual.pow(2)

Custom masking (mask override)

Override mask to exclude non-finite targets, padding, or other invalid entries. Return a boolean tensor broadcast-compatible with pred/target.

def mask(self, pred, target, ctx, **kwargs):
    if self.ignore_nonfinite:
        return torch.isfinite(target)
    return torch.ones_like(target, dtype=torch.bool)

For padded force layouts (B, V_max, 3), combine a node mask with nonfinite check:

def mask(self, pred, target, ctx, **kwargs):
    num_nodes_per_graph = kwargs.get("num_nodes_per_graph")
    node_mask = _padded_node_mask(num_nodes_per_graph, pred, pred.shape[1])
    valid = node_mask.unsqueeze(-1).expand_as(pred)
    if self.ignore_nonfinite:
        valid = valid & torch.isfinite(target)
    return valid

The valid tensor flows into compute_residual as the third argument. Zero invalid entries with torch.where(valid, ..., torch.zeros_like(...)).

Custom reduction (reduce override)

Override reduce for graph-balanced or other non-mean reductions. Populate self.per_sample_loss with a detached (B,) tensor for diagnostics.

from nvalchemi.training.losses.reductions import per_graph_sum

def reduce(self, residual, valid, ctx, **kwargs):
    batch_idx = kwargs["batch_idx"]
    num_graphs = kwargs["num_graphs"]
    valid_f = valid.to(dtype=residual.dtype)
    # Per-atom SE summed over xyz, then per-graph mean, then mean over graphs
    per_atom_se = residual.sum(dim=-1)
    per_atom_valid = valid_f.sum(dim=-1)
    per_graph_num = per_graph_sum(per_atom_se, batch_idx, num_graphs)
    per_graph_den = per_graph_sum(per_atom_valid, batch_idx, num_graphs)
    per_sample = per_graph_num / per_graph_den.clamp_min(1.0)
    self.per_sample_loss = per_sample.detach()
    return per_sample.mean()

Layout dispatch with plum (dense vs padded forces)

ForceMSELoss and ForceL2NormLoss use plum-dispatch to handle both dense (V, 3) and padded (B, V_max, 3) layouts without if/else on ndim. Their mask and reduce hooks delegate to @overload/@dispatch helper methods — one overload per layout. See these implementations in nvalchemi/training/losses/terms.py as the reference pattern for multi-layout losses.

from plum import dispatch, overload

@overload
def _my_helper(self, pred: Forces, target: Forces, ...):
    """Dense (V, 3) path."""
    ...

@overload
def _my_helper(self, pred: _PaddedForces, target: _PaddedForces, ...):
    """Padded (B, V_max, 3) path."""
    ...

@dispatch
def _my_helper(self, pred, target, ...):
    pass  # plum routes to matching overload at runtime

Conventions

  1. Define target_key and prediction_key on any loss that participates in ComposedLossFunction — these route tensors from the prediction/target mappings.
  2. Accept **kwargs in hooks that receive them — ComposedLossFunction forwards metadata kwargs to every component.
  3. compute_residual must zero invalid entries using the valid mask argument — the base reduce handles weighting but not masking.
  4. ReductionContext is a dict subclass (not TypedDict) for torch.compile compatibility. Conventional key: "weights" for atom-count weights consumed by the base reduce.

Key files

FileContents
nvalchemi/training/losses/composition.pyBaseLossFunction, ComposedLossFunction, ReductionContext
nvalchemi/training/losses/terms.pyAll 8 built-in leaf losses
nvalchemi/training/losses/reductions.pyper_graph_sum, per_graph_mean, frobenius_mse
nvalchemi/training/losses/schedules.pyConstantWeight, LinearWeight, CosineWeight, PiecewiseWeight
nvalchemi/training/losses/base.pyLossWeightSchedule protocol, re-exports
test/training/test_losses.pyComprehensive tests for all loss terms
docs/userguide/losses.mdFull user guide with examples

nvidia की और Skills

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 APIs कैसे काम करते हैं, कौन से संसाधन उपलब्ध हैं, उन्हें खोज मापदंडों के साथ कैसे क्वेरी करें, और सभी प्रतिक्रिया प्रारूपों को सही ढंग से कैसे पार्स करें…
compileiq-validate-result
nvidia
खोज पूरी होने के बाद और किसी स्पीडअप का दावा करने या ACF भेजने से पहले उपयोग करें। dump_results CSV लोड करता है, शीर्ष-K उम्मीदवारों (एकल-उद्देश्य) को निकालता है…
changelog-audit
nvidia
रिलीज़ से पहले Warp CHANGELOG.md का ऑडिट करें: खोई हुई प्रविष्टियाँ पुनर्प्राप्त करें, उपयोगकर्ता प्रभाव के अनुसार क्रमबद्ध करें, प्रविष्टि भाषा को परिष्कृत करें, लाइन-रैप करें, और (रिलीज़-ब्रांच मोड) तुलना बढ़ाएँ…
maintain-dynamic-plugins
nvidia
NeMo Relay डायनामिक प्लगइन लोडर, मैनिफेस्ट, रस्ट नेटिव SDK, gRPC वर्कर प्रोटोकॉल, पायथन वर्कर SDK, दस्तावेज़, परीक्षण और रिलीज़ वर्कफ़्लो कवरेज बनाए रखें
dgx-diagnose
nvidia
सामान्य DGX Station GB300 समस्याओं का निदान करें — CUDA क्रैश, गलत-GPU लक्ष्यीकरण, vLLM/SGLang कंटेनर बग, MIG स्थिति समस्याएं, NVLink/Fabric Manager त्रुटियां,…