File size: 4,857 Bytes
473253d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Reduction-to-discounted framework (Section 4, Theorem 4.2).

Solves the robust AMDP by reducing it to a robust DMDP with discount
gamma_reduction = 1 - eps / H, then solving the discounted problem.  Theorem 4.2
guarantees an eps_gamma-optimal DMDP policy is O(eps)-optimal for the AMDP.

This module provides:
  * the model-based reduction (exact discounted VI at gamma_reduction) to verify
    Theorem 4.2's policy-transfer guarantee numerically;
  * a model-free discounted robust solver (sample-based) to MEASURE the
    reduction's sample complexity and confirm it is worse (Õ(SA H^2/eps^4)) than
    RHI's direct approach (Õ(SA H^2/eps^2));
  * the algebraic derivation of the Õ(SA H^2/eps^4) bound.
"""
from __future__ import annotations

import numpy as np

from . import mdp as M


def reduction_gamma(eps: float, H: float) -> float:
    """gamma_reduction = 1 - eps / H  (Theorem 4.2).  REQUIRES knowledge of H."""
    return 1.0 - eps / H


def reduction_model_based(r: np.ndarray, P: np.ndarray, unc: M.Uncertainty,
                          eps: float, H: float, iters: int = 5000
                          ) -> tuple[np.ndarray, float, float]:
    """Exact discounted-VI reduction. Returns (pi, gamma, discounted subopt).

    Verifies Theorem 4.2: an (near-)optimal policy for the DMDP at gamma_reduction
    is O(eps)-optimal for the AMDP.
    """
    gamma = reduction_gamma(eps, H)
    V, Q = M.robust_discounted_vi(r, P, unc, gamma, iters=iters)
    pi = Q.argmax(axis=1)
    g_star, _, _ = M.robust_optimal_span_rvi(r, P, unc)
    g_pi, _ = M.robust_policy_gain(pi, r, P, unc)
    return pi, gamma, g_star - g_pi


def discounted_sample_based(r: np.ndarray, P0: np.ndarray, unc: M.Uncertainty,
                            gamma: float, n_samples: int,
                            rng: np.random.Generator, iters: int = 300
                            ) -> tuple[np.ndarray, int]:
    """Model-free robust discounted VI: estimate sigma from n_samples per (s,a).

    A faithful generative-model discounted robust solver.  Returns (pi, budget)
    where budget = S*A*n_samples*iters is the total transition samples consumed.
    """
    S, A = r.shape
    V = np.zeros(S)
    for _ in range(iters):
        # estimate sigma_{P_sa}(V) from n_samples for every (s,a)
        sig = np.empty((S, A))
        for s in range(S):
            for a in range(A):
                Prow = P0[s, a]
                idx = rng.choice(len(Prow), size=n_samples, p=Prow)
                R = unc.radius if np.isscalar(unc.radius) else unc.radius[s, a]
                emp = float(np.mean(V[idx]))
                if unc.kind == "contamination":
                    sig[s, a] = (1 - R) * emp + R * float(np.min(V))
                else:
                    reach = Prow > 0
                    sig[s, a] = emp - R * M.kappa_q(V, unc.p, reach)
        Q = r + gamma * sig
        Vn = Q.max(axis=1)
        if np.max(np.abs(Vn - V)) < 1e-7:
            V = Vn
            break
        V = Vn
    pi = (r + gamma * _sigma_from_V(V, P0, unc)).argmax(axis=1)
    budget = S * A * n_samples * iters
    return pi, budget


def _sigma_from_V(V, P0, unc):
    S, A = P0.shape[:2]
    out = np.empty((S, A))
    for s in range(S):
        for a in range(A):
            Prow = P0[s, a]
            R = unc.radius if np.isscalar(unc.radius) else unc.radius[s, a]
            reach = Prow > 0
            if unc.kind == "contamination":
                out[s, a] = (1 - R) * np.dot(Prow, V) + R * np.min(V)
            else:
                out[s, a] = np.dot(Prow, V) - R * M.kappa_q(V, unc.p, reach)
    return out


# --- algebraic derivation of the Õ(SA H^2 / eps^4) bound ---------------------

def reduction_algebra(eps: float, H: float, S: int, A: int) -> dict:
    """Reconstruct the Õ(SA H^2 / eps^4) sample-complexity derivation (Sec. 4).

    robust DMDP sample complexity (Clavier et al. 2024, lp) with discount gamma
    and target eps_gamma:  N = Õ( SA / ((1-gamma) * eps_gamma^2) )  using the
    effective-horizon 1/(1-gamma).  Reduction sets gamma = 1 - eps/H and needs
    eps_gamma = Theta(eps * (1-gamma)) = Theta(eps^2/H) for the discounted value
    to be eps-optimal after the (1-gamma) rescaling -> N = Õ(SA H^2 / eps^4).

    We expose each factor so the derivation is machine-checkable, and also quote
    the paper's direct statement (eps_gamma = H) which gives the same H^2/eps^4.
    """
    gamma = reduction_gamma(eps, H)
    one_minus = eps / H
    # paper's direct choice: eps_gamma = H
    eps_gamma_paper = H
    N_paper = float(S * A / (one_minus ** 2 * eps_gamma_paper ** 2))  # -> SA H^2/eps^4
    return dict(eps=eps, H=H, S=S, A=A, gamma=gamma, one_minus_gamma=one_minus,
                eps_gamma_paper=eps_gamma_paper,
                N_paper_tilde=N_paper,
                expected_H_power=2, expected_eps_power=4)