|
| 1 | +import type { HorizonConfig } from './config.js'; |
| 2 | +import type { |
| 3 | + HorizonRollout, |
| 4 | + SnapshotStateFn, |
| 5 | + HorizonStep, |
| 6 | + HorizonStepRole, |
| 7 | + HorizonTransitionReason, |
| 8 | +} from './types.js'; |
| 9 | +import type { PolicyEvaluator } from './policy.js'; |
| 10 | +import { choosePlanningAction } from './policy.js'; |
| 11 | +import type { TransitionFn } from './transition.js'; |
| 12 | +import { transitionFutureState } from './transition.js'; |
| 13 | + |
| 14 | +function stepRoleFor(step: number): HorizonStepRole { |
| 15 | + return step === 0 ? 'selected_present_action' : 'projected_future_action'; |
| 16 | +} |
| 17 | + |
| 18 | +function transitionReasonFor<TAction>( |
| 19 | + stepRole: HorizonStepRole, |
| 20 | + action: TAction | null |
| 21 | +): HorizonTransitionReason { |
| 22 | + if (action === null) { |
| 23 | + return 'projected_state_only'; |
| 24 | + } |
| 25 | + |
| 26 | + return stepRole === 'selected_present_action' |
| 27 | + ? 'selected_present_action_projected' |
| 28 | + : 'projected_future_action'; |
| 29 | +} |
| 30 | + |
| 31 | +const defaultSnapshotState = <TState>(state: TState): TState => |
| 32 | + structuredClone(state); |
| 33 | + |
| 34 | +export function runHorizonRollout<TState, TAction, TObjective>(args: { |
| 35 | + initialState: TState; |
| 36 | + config: HorizonConfig; |
| 37 | + evaluatePolicy: PolicyEvaluator<TState, TAction, TObjective>; |
| 38 | + applyAction: TransitionFn<TState, TAction>; |
| 39 | + snapshotState?: SnapshotStateFn<TState>; |
| 40 | +}): HorizonRollout<TState, TAction, TObjective> { |
| 41 | + const steps: HorizonStep<TState, TAction, TObjective>[] = []; |
| 42 | + const snapshotState = |
| 43 | + args.snapshotState === undefined ? defaultSnapshotState : args.snapshotState; |
| 44 | + |
| 45 | + let state = args.initialState; |
| 46 | + let selectedPresentAction: TAction | null = null; |
| 47 | + |
| 48 | + for (let step = 0; step < args.config.steps; step += 1) { |
| 49 | + const stateBefore = snapshotState(state); |
| 50 | + const stepRole = stepRoleFor(step); |
| 51 | + const { action, objective } = choosePlanningAction({ |
| 52 | + state, |
| 53 | + step, |
| 54 | + evaluatePolicy: args.evaluatePolicy, |
| 55 | + }); |
| 56 | + |
| 57 | + if (step === 0) { |
| 58 | + selectedPresentAction = action; |
| 59 | + } |
| 60 | + |
| 61 | + const stateAfter = |
| 62 | + action === null |
| 63 | + ? state |
| 64 | + : transitionFutureState({ |
| 65 | + state, |
| 66 | + action, |
| 67 | + step, |
| 68 | + applyAction: args.applyAction, |
| 69 | + }); |
| 70 | + const stateAfterSnapshot = snapshotState(stateAfter); |
| 71 | + |
| 72 | + steps.push({ |
| 73 | + step, |
| 74 | + label: 'planning_projection', |
| 75 | + stepRole, |
| 76 | + stateBefore, |
| 77 | + action, |
| 78 | + objective, |
| 79 | + stateAfter: stateAfterSnapshot, |
| 80 | + transitionReason: transitionReasonFor(stepRole, action), |
| 81 | + }); |
| 82 | + |
| 83 | + state = stateAfter; |
| 84 | + } |
| 85 | + |
| 86 | + return { |
| 87 | + label: 'planning_projection', |
| 88 | + horizonSteps: args.config.steps, |
| 89 | + futureJustification: args.config.futureJustification, |
| 90 | + selectedPresentAction, |
| 91 | + steps, |
| 92 | + }; |
| 93 | +} |
0 commit comments