import * as ort from "onnxruntime-web/all"; import { currentHandSorted } from "../game/engine"; import { matchLegalActionMask, type MatchState } from "../game/match"; import { matchObservation } from "../game/matchObservation"; import { DISCARD, DRAW_DECK, PLAY } from "../game/types"; import { cardColor, cardRank, isHandshake } from "../game/cards"; export type ExecutionProvider = "webgpu" | "wasm" | "heuristic"; /** * borealis -- see data/models.json. Trained on the three-round match, so it takes * the match view (carry, round, deck clock) rather than a bare round. A one-off * deal is simply its round one at a carry of zero, which is a position it has seen * a great many times. */ const MODEL_URL = `${import.meta.env.BASE_URL}models/borealis.onnx`; /** Identity of what actually plays. Records carry the hash; the codename is for * humans and is assigned in data/models.json, not derived. */ export const MODEL_CODENAME = "borealis"; export const MODEL_HASH = "13a25243de1c"; /** A legal action with the policy's confidence in it, if the policy has one. */ export interface RankedAction { action: number; probability: number | null; } export interface Policy { readonly provider: ExecutionProvider; /** Legal actions, best first. Drives both the rival's move and the hint. */ rank(match: MatchState): Promise; action(match: MatchState): Promise; } abstract class RankingPolicy implements Policy { abstract readonly provider: ExecutionProvider; abstract rank(match: MatchState): Promise; async action(match: MatchState): Promise { const [best] = await this.rank(match); if (best === undefined) throw new Error("state has no legal actions"); return best.action; } } /** Softmax over the legal actions only, so the reported confidence is a share * of what the policy could actually have chosen. */ function softmaxOverLegal(scores: number[], legal: boolean[]): RankedAction[] { const legalScores = scores.filter((_, action) => legal[action]); const max = Math.max(...legalScores); const total = legalScores.reduce((sum, score) => sum + Math.exp(score - max), 0); return scores .flatMap((score, action) => legal[action] ? [{ action, probability: Math.exp(score - max) / total }] : [], ) .sort((left, right) => right.probability - left.probability); } class OnnxPolicy extends RankingPolicy { constructor( private readonly session: ort.InferenceSession, readonly provider: ExecutionProvider, ) { super(); } async rank(match: MatchState): Promise { const obs = matchObservation(match, match.round.toMove); const result = await this.session.run({ obs: new ort.Tensor("float32", obs, [1, obs.length]) }); const logits = result.logits?.data; if (!logits) throw new Error("ONNX model did not return a logits output"); const legal = matchLegalActionMask(match); return softmaxOverLegal(legal.map((_, action) => Number(logits[action])), legal); } } export class HeuristicPolicy extends RankingPolicy { readonly provider = "heuristic" as const; async rank(match: MatchState): Promise { const state = match.round; const legal = matchLegalActionMask(match); const hand = currentHandSorted(state); const ranked = legal.flatMap((isLegal, action) => { if (!isLegal) return []; const handSlot = Math.floor(action / 12); const place = Math.floor((action % 12) / 6); const draw = action % 6; const card = hand[handSlot]; let score = place === PLAY ? cardRank(card) + (isHandshake(card) ? 7 : 0) : -cardRank(card); if (draw !== DRAW_DECK) score += 2; if (place === DISCARD && draw > 0 && draw - 1 === cardColor(card)) score -= 100; return [{ action, score }]; }); // No calibrated probability to report — this is a hand-written score. return ranked .sort((left, right) => right.score - left.score) .map(({ action }) => ({ action, probability: null })); } } async function createSession(provider: "webgpu" | "wasm"): Promise { return ort.InferenceSession.create(MODEL_URL, { executionProviders: [provider], graphOptimizationLevel: "all", }); } let policyPromise: Promise<{ policy: Policy; warning?: string }> | undefined; async function initializePolicy(): Promise<{ policy: Policy; warning?: string }> { const canUseWebGpu = "gpu" in navigator; if (canUseWebGpu) { try { return { policy: new OnnxPolicy(await createSession("webgpu"), "webgpu") }; } catch (error) { console.warn("WebGPU model initialization failed; trying WASM", error); } } try { return { policy: new OnnxPolicy(await createSession("wasm"), "wasm") }; } catch (error) { console.warn("ONNX model initialization failed; using heuristic policy", error); return { policy: new HeuristicPolicy(), warning: "ONNX 모델을 찾지 못해 로컬 휴리스틱으로 플레이합니다.", }; } } export function loadPolicy(): Promise<{ policy: Policy; warning?: string }> { policyPromise ??= initializePolicy(); return policyPromise; } let heuristicFallback: HeuristicPolicy | undefined; export function fallbackHeuristicPolicy(): HeuristicPolicy { heuristicFallback ??= new HeuristicPolicy(); return heuristicFallback; }