"""
行为克隆 vs 朴素离线 Q 的直觉对比（阶段 3 配套）
------------------------------------------------
场景：1 维状态格子 0..N-1，动作 Left/Right，只有历史日志，不能再交互。

1) Behavior Cloning：模仿日志里的动作（监督）
2) 朴素离线 Q-learning：对未见动作做 max_a Q → 容易对 OOD 动作过度乐观

运行：python bc_offline_demo.py
"""

from __future__ import annotations

import random
from collections import defaultdict

N = 7          # 状态 0..6，目标 6
ACTIONS = (-1, 1)  # Left, Right
GAMMA = 0.95
ALPHA = 0.3
EPISODES_DATA = 40
SEED = 0


def step(s: int, a: int) -> tuple[int, float, bool]:
    ns = max(0, min(N - 1, s + a))
    if ns == N - 1:
        return ns, 1.0, True
    return ns, -0.05, False


def collect_dataset(rng: random.Random) -> list[tuple[int, int, float, int, bool]]:
    """
    行为策略：偏右但偶尔乱走。只记录轨迹，之后不再与环境交互训练 offline Q。
    """
    data = []
    for _ in range(EPISODES_DATA):
        s = 0
        for _ in range(30):
            a = 1 if rng.random() < 0.75 else rng.choice(ACTIONS)
            ns, r, done = step(s, a)
            data.append((s, a, r, ns, done))
            s = ns
            if done:
                break
    return data


def behavior_clone(data):
    """BC：每个状态选日志里最常见动作。"""
    counts = defaultdict(lambda: defaultdict(int))
    for s, a, _, _, _ in data:
        counts[s][a] += 1
    pi = {}
    for s in range(N):
        if counts[s]:
            pi[s] = max(counts[s].items(), key=lambda kv: kv[1])[0]
        else:
            pi[s] = 1  # 默认向右
    return pi, counts


def naive_offline_q(data, iters=200):
    """
    朴素 Q-learning 风格更新（离线、无保守项）。
    对每个状态 max_a Q(s,a) —— 即便 a 几乎没在数据中出现。
    """
    Q = {(s, a): 0.0 for s in range(N) for a in ACTIONS}
    for _ in range(iters):
        for s, a, r, ns, done in data:
            target = r if done else r + GAMMA * max(Q[(ns, a2)] for a2 in ACTIONS)
            Q[(s, a)] += ALPHA * (target - Q[(s, a)])
    pi = {}
    for s in range(N):
        pi[s] = max(ACTIONS, key=lambda a: Q[(s, a)])
    return Q, pi


def action_coverage(data):
    cov = defaultdict(set)
    for s, a, _, _, _ in data:
        cov[s].add(a)
    return cov


def rollout(pi, max_steps=30) -> float:
    s = 0
    total = 0.0
    for _ in range(max_steps):
        a = pi.get(s, 1)
        s, r, done = step(s, a)
        total += r
        if done:
            break
    return total


def main():
    rng = random.Random(SEED)
    data = collect_dataset(rng)
    cov = action_coverage(data)

    pi_bc, counts = behavior_clone(data)
    Q, pi_q = naive_offline_q(data)

    print("=== BC vs 朴素离线 Q ===")
    print(f"离线转移条数: {len(data)}")
    print("\n各状态动作覆盖（数据里真实出现过的 a）:")
    for s in range(N):
        print(f"  s={s}: {sorted(cov[s])}  counts={dict(counts[s])}")

    print("\n学到的策略（-1=Left, +1=Right）:")
    print("  state:", list(range(N)))
    print("  BC   :", [pi_bc[s] for s in range(N)])
    print("  Q-pol:", [pi_q[s] for s in range(N)])

    print("\nQ 表（注意：覆盖差的 (s,a) 也可能被 max 抬很高）:")
    for s in range(N):
        row = {a: round(Q[(s, a)], 3) for a in ACTIONS}
        print(f"  s={s}: {row}  covered={sorted(cov[s])}")

    # 在线评估（仅评估，不用于训练）
    rets_bc = [rollout(pi_bc) for _ in range(20)]
    rets_q = [rollout(pi_q) for _ in range(20)]
    print("\n在线评估平均回报（训练时不可用）:")
    print(f"  BC:           {sum(rets_bc)/len(rets_bc):.3f}")
    print(f"  朴素离线 Q:   {sum(rets_q)/len(rets_q):.3f}")

    print("\n解读:")
    print("  - BC 只模仿数据里常见动作，通常更稳，但上限受专家限制。")
    print("  - 朴素离线 max-Q 会对覆盖不足的动作过度乐观 → 策略可能选 OOD 动作。")
    print("  - 真正的离线 RL（CQL/IQL 等）要加保守项：压低「数据少见动作」的 Q。")
    print("  - 生物实验日志：覆盖窄时先 BC；有多样失败/成功与回报再考虑离线 RL。")


if __name__ == "__main__":
    main()
