File size: 3,633 Bytes
5ccb4fd | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 | # Copyright (c) 2026 Simulacra Research Inc.
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
import jax
import jax.numpy as jnp
_ROUTE_SAMPLES = 8
def _clip_one_channel(x, width: float):
median = jnp.nanmedian(x, axis=1, keepdims=True)
mean_ad = jnp.nanmean(jnp.abs(x - median), axis=1, keepdims=True)
delta = jnp.asarray(width, dtype=x.dtype) * mean_ad
return jnp.clip(x, median - delta, median + delta)
def _mad_clip_per_system(local_energies, width: float):
re = jnp.real(local_energies)
im = jnp.imag(local_energies)
re_clipped = _clip_one_channel(re, width)
im_clipped = _clip_one_channel(im, width)
if jnp.iscomplexobj(local_energies):
clipped = (re_clipped + 1j * im_clipped).astype(local_energies.dtype)
else:
clipped = re_clipped.astype(local_energies.dtype)
return clipped
def process_route_targets(
sampled_energy,
baseline_energy,
sigma1,
baseline_weights,
*,
mad_width: float = 5.0,
):
systems = int(sampled_energy.shape[0])
if systems % _ROUTE_SAMPLES:
raise ValueError("router targets require a system axis divisible by K=8")
sigma = sigma1.astype(sampled_energy.real.dtype)
sampled_normalized = sampled_energy / sigma[:, None]
baseline_normalized = baseline_energy / sigma[:, None]
sampled_clipped = _mad_clip_per_system(
sampled_normalized,
mad_width,
)
centered = sampled_clipped - jnp.mean(
sampled_clipped,
axis=1,
keepdims=True,
)
variance = jnp.mean(
centered.real**2 + centered.imag**2,
axis=1,
keepdims=True,
)
group_variance = jnp.mean(
variance.reshape(systems // _ROUTE_SAMPLES, _ROUTE_SAMPLES),
axis=1,
keepdims=True,
)
group_std = jnp.sqrt(
jnp.broadcast_to(
group_variance,
(systems // _ROUTE_SAMPLES, _ROUTE_SAMPLES),
)
).reshape(systems, 1)
scale = jnp.maximum(group_std, jnp.asarray(1.0, dtype=group_std.dtype))
sampled_target = sampled_clipped / scale
baseline_target = baseline_normalized / scale
sampled_rewards = jnp.mean(sampled_target.real, axis=1).astype(jnp.float32)
baseline_rewards = jnp.sum(
baseline_weights * baseline_target.real,
axis=1,
).astype(jnp.float32)
grouped_baseline = baseline_rewards.reshape((-1, _ROUTE_SAMPLES))
group_is_finite = jnp.all(
jnp.isfinite(grouped_baseline),
axis=1,
keepdims=True,
)
baseline_rewards = jnp.where(
group_is_finite,
grouped_baseline,
jnp.zeros_like(grouped_baseline),
).reshape(baseline_rewards.shape)
reward_delta = sampled_rewards.reshape(
(-1, _ROUTE_SAMPLES)
) - baseline_rewards.reshape((-1, _ROUTE_SAMPLES))
advantage = (
float(_ROUTE_SAMPLES)
/ float(_ROUTE_SAMPLES - 1)
* (reward_delta - jnp.mean(reward_delta, axis=1, keepdims=True))
)
advantage = jax.lax.stop_gradient(
advantage.reshape(sampled_rewards.shape).astype(jnp.float32)
)
return sampled_target, advantage
def process_finetune_targets(
energy,
sigma1,
*,
mad_width: float = 5.0,
):
normalized = energy / sigma1.astype(energy.real.dtype)[:, None]
clipped = _mad_clip_per_system(normalized, mad_width)
centered = clipped - jnp.mean(clipped, axis=1, keepdims=True)
std = jnp.sqrt(jnp.mean(centered.real**2 + centered.imag**2, axis=1))
return clipped / jnp.maximum(std, 1.0)[:, None]
__all__ = ["process_finetune_targets", "process_route_targets"]
|