File size: 10,840 Bytes
23a59ea
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
# task_set.py

# 200 pretraining tasks (the 10 held-out unseen tasks are in UNSEEN_TASK_SET below)
TASK_SET = [
    # dmcontrol (21 tasks)
    'walker-stand', 'walker-walk', 'walker-run', 'cheetah-run', 'reacher-easy',
    'reacher-hard', 'acrobot-swingup', 'pendulum-swingup', 'cartpole-balance', 'cartpole-balance-sparse',
    'cartpole-swingup', 'cartpole-swingup-sparse', 'cup-catch', 'finger-spin', 'finger-turn-easy',
    'finger-turn-hard', 'fish-swim', 'hopper-stand', 'hopper-hop', 'quadruped-walk',
    'quadruped-run',
    # dmcontrol-ext (16 tasks)
    'walker-walk-backward', 'walker-run-backward', 'cheetah-run-backward', 'cheetah-run-front', 'cheetah-run-back',
    'cheetah-jump', 'hopper-hop-backward', 'reacher-three-easy', 'reacher-three-hard', 'cup-spin',
    'pendulum-spin', 'jumper-jump', 'spinner-spin', 'spinner-spin-backward', 'spinner-jump',
    'giraffe-run',
    # meta-world (49 tasks)
    'mw-assembly', 'mw-basketball', 'mw-button-press-topdown', 'mw-button-press-topdown-wall', 'mw-button-press',
    'mw-button-press-wall', 'mw-coffee-button', 'mw-coffee-pull', 'mw-coffee-push', 'mw-dial-turn',
    'mw-disassemble', 'mw-door-open', 'mw-door-close', 'mw-drawer-close', 'mw-drawer-open',
    'mw-faucet-open', 'mw-faucet-close', 'mw-hammer', 'mw-handle-press-side', 'mw-handle-press',
    'mw-handle-pull-side', 'mw-handle-pull', 'mw-lever-pull', 'mw-peg-insert-side', 'mw-peg-unplug-side',
    'mw-pick-out-of-hole', 'mw-pick-place', 'mw-pick-place-wall', 'mw-plate-slide', 'mw-plate-slide-side',
    'mw-plate-slide-back', 'mw-plate-slide-back-side', 'mw-push-back', 'mw-push', 'mw-push-wall',
    'mw-reach', 'mw-reach-wall', 'mw-soccer', 'mw-stick-push', 'mw-stick-pull', 
    'mw-sweep-into', 'mw-sweep', 'mw-window-open', 'mw-window-close', 'mw-bin-picking',
    'mw-box-close', 'mw-door-lock', 'mw-door-unlock', 'mw-hand-insert',
    # maniskill (36 tasks)
    'ms-ant-walk', 'ms-ant-run', 'ms-cartpole-balance', 'ms-cartpole-swingup', 'ms-hopper-stand',
    'ms-hopper-hop', 'ms-pick-cube', 'ms-pick-cube-eepose', 'ms-pick-cube-so', 'ms-poke-cube',
    'ms-push-cube', 'ms-pull-cube', 'ms-pull-cube-tool', 'ms-stack-cube', 'ms-place-sphere',
    'ms-lift-peg', 'ms-pick-apple', 'ms-pick-banana', 'ms-pick-can', 'ms-pick-hammer',
    'ms-pick-fork', 'ms-pick-knife', 'ms-pick-mug', 'ms-pick-orange', 'ms-pick-screwdriver',
    'ms-pick-spoon', 'ms-pick-tennis-ball', 'ms-pick-baseball', 'ms-pick-cube-xarm6', 'ms-pick-sponge',
    'ms-anymal-reach', 'ms-reach', 'ms-reach-eepose', 'ms-reach-xarm6', 'ms-cartpole-balance-sparse',
    'ms-cartpole-swingup-sparse',
    # mujoco (6 tasks)
    'mujoco-ant', 'mujoco-halfcheetah', 'mujoco-hopper', 'mujoco-inverted-pendulum', 'mujoco-reacher',
    'mujoco-walker',
    # box2d (8 tasks)
    'bipedal-walker-flat', 'bipedal-walker-uneven', 'bipedal-walker-rugged', 'bipedal-walker-hills', 'bipedal-walker-obstacles',
    'lunarlander-land', 'lunarlander-hover', 'lunarlander-takeoff',
    # robodesk (6 tasks)
    'rd-push-red', 'rd-push-green', 'rd-push-blue', 'rd-open-slide', 'rd-open-drawer',
    'rd-flat-block-in-bin',
    # ogbench (12 tasks)
    'og-ant', 'og-antball', 'og-point-arena', 'og-point-maze', 'og-point-bottleneck',
    'og-point-circle', 'og-point-spiral', 'og-ant-arena', 'og-ant-maze', 'og-ant-bottleneck',
    'og-ant-circle', 'og-ant-spiral',
    # pygame (19 tasks)
    'pygame-cowboy', 'pygame-coinrun', 'pygame-spaceship', 'pygame-pong', 'pygame-bird-attack',
    'pygame-highway', 'pygame-landing', 'pygame-air-hockey', 'pygame-rocket-collect', 'pygame-chase-evade',
    'pygame-coconut-dodge', 'pygame-cartpole-balance', 'pygame-cartpole-swingup', 'pygame-cartpole-balance-sparse', 'pygame-cartpole-swingup-sparse',
    'pygame-cartpole-tremor', 'pygame-point-maze-var1', 'pygame-point-maze-var2', 'pygame-point-maze-var3',
    # atari (27 tasks)
    'atari-alien', 'atari-assault', 'atari-asterix', 'atari-atlantis', 'atari-bank-heist',
    'atari-battle-zone', 'atari-beamrider', 'atari-boxing', 'atari-chopper-command', 'atari-crazy-climber',
    'atari-double-dunk', 'atari-gopher', 'atari-ice-hockey', 'atari-jamesbond', 'atari-kangaroo',
    'atari-krull', 'atari-ms-pacman', 'atari-name-this-game', 'atari-phoenix', 'atari-pong',
    'atari-road-runner', 'atari-robotank', 'atari-seaquest', 'atari-space-invaders', 'atari-tutankham',
    'atari-upndown', 'atari-yars-revenge',
]


# Domain spans in TASK_SET ordering.  Task-set authoritatively defines the
# 10 domains below; keep in sync with the task blocks above.
DOMAIN_SPANS: list[tuple[str, int, int]] = [
    ('dmcontrol',      0,  21),
    ('dmcontrol-ext', 21,  37),
    ('metaworld',     37,  86),
    ('maniskill',     86, 122),
    ('mujoco',       122, 128),
    ('box2d',        128, 136),
    ('robodesk',     136, 142),
    ('ogbench',      142, 154),
    ('pygame',       154, 173),
    ('atari',        173, 200),
]
DOMAINS: list[str] = [name for name, _, _ in DOMAIN_SPANS]

_TASK_TO_DOMAIN: dict[str, str] = {}
for _name, _lo, _hi in DOMAIN_SPANS:
    for _t in TASK_SET[_lo:_hi]:
        _TASK_TO_DOMAIN[_t] = _name


def task_to_domain(task: str) -> str:
    """Map a task name to its canonical domain string.

    Covers TASK_SET ∪ UNSEEN_TASK_SET. Raises KeyError otherwise — the
    caller should fall back to a default domain or update the table if
    this fires.
    """
    return _TASK_TO_DOMAIN[task]


# ---------------------------------------------------------------------------
# Seen / unseen task sets for the targeted-data-collection evaluation.
#   - SEEN_TASK_SET: subset of TASK_SET (tasks the world model trained on).
#   - UNSEEN_TASK_SET: tasks defined in envs/*.py but NOT in TASK_SET (also
#     enumerated as keys in interactive.TEST_TASK_SET).
# Each matched UNSEEN entry is the in-domain analog of a SEEN entry (same
# domain / control morphology), since out-of-domain tasks transfer poorly
# zero-shot.
#
# IMPORTANT: three matched pairs share IDENTICAL CLIP text embeddings
# (cosine = 1.0000) by design: cup-catch ↔ cup-catch-var1, finger-turn-easy ↔
# finger-turn-easy-var1, pygame-point-maze-var1 ↔ pygame-point-maze-var4. The
# model thus has no language signal to disambiguate them, so any UNSEEN-partner
# gain is attributable to visual adaptation, not language conditioning (and
# such pairs are not "true zero-shot" in the language sense). The three pygame
# tasks dungeon-explorer1/foraging/whirlpool have no SEEN partner.
# ---------------------------------------------------------------------------

SEEN_TASK_SET: list[str] = [
    'cup-catch',                # dmcontrol
    'finger-turn-easy',         # dmcontrol
    'mw-push',                  # metaworld
    'ms-push-cube',             # maniskill
    'lunarlander-hover',        # box2d
    'og-point-maze',            # ogbench
    'og-point-bottleneck',      # ogbench
    'pygame-point-maze-var1',   # pygame
    'pygame-pong',              # pygame
    'pygame-bird-attack',       # pygame
]

UNSEEN_TASK_SET: list[str] = [
    # Matched-pair UNSEEN: each entry has a SEEN partner with shared visual style.
    'cup-catch-var1',           # dmcontrol; SEEN partner: cup-catch        (lang cosine 1.0000)
    'finger-turn-easy-var1',    # dmcontrol; SEEN partner: finger-turn-easy (lang cosine 1.0000)
    'ms-push-banana',           # maniskill; SEEN partner: ms-push-cube     (lang cosine 0.986)
    'og-point-var1',            # ogbench;   SEEN partner: og-point-bottleneck (lang cosine 0.992)
    'og-point-var2',            # ogbench;   SEEN partner: og-point-maze       (lang cosine 0.994)
    'pygame-point-maze-var4',   # pygame;    SEEN partner: pygame-point-maze-var1 (lang cosine 1.0000)
    'pygame-reacher-easy',      # pygame;    SEEN partner: reacher-easy (dmcontrol; same instruction)
    # Completely unseen: no SEEN partner.
    'pygame-dungeon-explorer1',
    'pygame-foraging',
    'pygame-whirlpool',
]

# UNSEEN tasks are not part of the contiguous DOMAIN_SPANS over TASK_SET, so
# extend the lookup table explicitly (matched-per-domain home per the
# inline comments above).
_UNSEEN_TASK_TO_DOMAIN: dict[str, str] = {
    'cup-catch-var1':           'dmcontrol',
    'finger-turn-easy-var1':    'dmcontrol',
    'ms-push-banana':           'maniskill',
    'og-point-var1':            'ogbench',
    'og-point-var2':            'ogbench',
    'pygame-point-maze-var4':   'pygame',
    'pygame-reacher-easy':      'pygame',
    'pygame-dungeon-explorer1': 'pygame',
    'pygame-foraging':          'pygame',
    'pygame-whirlpool':         'pygame',
}
_TASK_TO_DOMAIN.update(_UNSEEN_TASK_TO_DOMAIN)


def compute_task_weights(
    tasks: list[str],
    mode: str,
    *,
    targeted_alpha: float = 0.5,
    targeted_tasks: list[str] | None = None,
) -> list[float]:
    """Build per-task sampling weights for WMDataset.

    Modes (all weights are relative; WMDataset normalizes internally):
      - "valid_starts": returns None-equivalent (signals caller to pass
        task_weights=None for legacy P(task) ∝ total valid_starts).
      - "uniform": every task gets weight 1.0 (equal-per-task sampling).
      - "targeted": split total weight `α` uniformly over the focused task
        set and `1-α` uniformly over the rest, biasing sampling toward the
        targeted tasks while retaining the rest. Defaults: `targeted_alpha=0.5`,
        `targeted_tasks=SEEN_TASK_SET ∪ UNSEEN_TASK_SET`.

    The `tasks` list is what the dataset actually loaded (post filter);
    unknown tasks get weight 1.0 under "uniform" (shouldn't happen with
    the canonical TASK_SET).
    """
    if mode == "valid_starts":
        # Sentinel: caller should pass task_weights=None.
        return [1.0] * len(tasks)
    if mode == "uniform":
        return [1.0] * len(tasks)
    if mode == "targeted":
        if not (0.0 <= targeted_alpha <= 1.0):
            raise ValueError(
                f"targeted_alpha must be in [0,1], got {targeted_alpha!r}"
            )
        targeted_set = set(SEEN_TASK_SET + UNSEEN_TASK_SET if targeted_tasks is None else targeted_tasks)
        n_targeted = sum(1 for t in tasks if t in targeted_set)
        n_non = len(tasks) - n_targeted
        # Auto-collapse alpha when one side is empty so we degrade
        # gracefully to "uniform over what's loaded" instead of all-zero
        # weights (which would crash WMDataset's sampler).
        if n_targeted == 0:
            alpha_eff = 0.0
        elif n_non == 0:
            alpha_eff = 1.0
        else:
            alpha_eff = float(targeted_alpha)
        w_targeted = (alpha_eff / n_targeted) if n_targeted > 0 else 0.0
        w_non = ((1.0 - alpha_eff) / n_non) if n_non > 0 else 0.0
        return [w_targeted if t in targeted_set else w_non for t in tasks]
    raise ValueError(f"Unknown task_weighting mode: {mode!r}")