File size: 2,313 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
# Copyright (c) 2026 Simulacra Research Inc.
# SPDX-License-Identifier: Apache-2.0

from __future__ import annotations

import equinox as eqx
import jax.numpy as jnp

from hamiltonzero.model.context import MultiSystemContext, SpinContext
from .permutation import permute_ctx_prefix, permute_multi_ctx_prefix, permute_q_prefix


def batch_context(context: SpinContext) -> MultiSystemContext:
    return MultiSystemContext.from_single(context)


def select_frozen_route(
    model,
    context: SpinContext,
    *,
    tau: float,
):
    node, edge, global_features = model.route_features(context)
    permutations, _log_probabilities = model.route_decoder.beam_search(
        node,
        edge,
        context.bmask,
        global_feat=global_features,
        tau=tau,
        beam_width=8,
        real_mask=context.mask,
        first_orbit_ids=(
            context.route_quotient_node_key,
            context.route_quotient_edge_key,
        ),
    )
    return permutations[0].astype(jnp.int32)


def route_context(context, perm):
    if isinstance(context, MultiSystemContext):
        perms = perm if perm.ndim == 2 else perm[None, :]
        routed = permute_multi_ctx_prefix(context, perms)
        return eqx.tree_at(lambda c: c.route_perm, routed, perms)
    routed = permute_ctx_prefix(context, perm)
    return eqx.tree_at(lambda c: c.route_perm, routed, perm)


def route_state(state, perm):
    if state.q.ndim == 5:
        perms = perm if perm.ndim == 2 else perm[None, :]
        q = permute_q_prefix(state.q, perms)
        grad = permute_q_prefix(state.grad_log_p, perms)
        mask = jnp.take_along_axis(state.mask, perms, axis=-1)
    else:
        p = perm if perm.ndim == 1 else perm[0]
        q = permute_q_prefix(state.q[None, ...], p[None, :])[0]
        grad = permute_q_prefix(state.grad_log_p[None, ...], p[None, :])[0]
        mask = jnp.take_along_axis(state.mask, p, axis=-1)
    return eqx.tree_at(
        lambda s: (s.q, s.grad_log_p, s.mask),
        state,
        (q, grad, mask),
    )


def strip_router(model):
    return eqx.tree_at(
        lambda m: (m.route_decoder, m.route_contextualizer),
        model,
        (None, None),
    )


__all__ = [
    "batch_context",
    "route_context",
    "route_state",
    "select_frozen_route",
    "strip_router",
]