|
| 1 | +""" |
| 2 | +Provisioning test for the female mate-proximity association (RESULTS.md, Iteration 14). |
| 3 | +
|
| 4 | +Hypothesis under test: a female approaches food less when her recorded mate is nearby because he can donate part |
| 5 | +of a successful hunt to her (`_apply_male_gift`, only to his recorded mate within predator_gift_range), so a |
| 6 | +nearby, recently-successful mate means she is being fed and forages less urgently. This is a first, sharper |
| 7 | +test than a bare correlation of per-life gifts with approach bias (which would also pick up survival duration, |
| 8 | +location, the mate's own success, and so on). Two contrasts, females only, in the env each run was trained in: |
| 9 | +
|
| 10 | + (1) Gift timing. Females with a living mate within --near-radius are split by whether they RECEIVED a gift |
| 11 | + within the last --gift-window steps ("gift_recent") or not ("near_nogift"). If provisioning drives the |
| 12 | + association, near_nogift should look like "away" (far + dead + abandoned) and gift_recent should carry |
| 13 | + the drop. Reported: near_nogift-away, gift_recent-away, gift_recent-near_nogift. |
| 14 | + (2) Energy adjustment. The female's own energy (visible to the policy) is a natural mediator: a fed female is |
| 15 | + richer. The near-minus-away approach-bias difference is recomputed WITHIN bands of her own energy and |
| 16 | + averaged with weights equal to each band's pooled decision count (direct standardization). If energy carries |
| 17 | + the association, the adjusted difference shrinks toward zero relative to the crude one. |
| 18 | +
|
| 19 | +All intervals are 95% episode-cluster bootstrap percentiles (whole episodes resampled; every statistic is |
| 20 | +recomputed inside each replicate). What this can and cannot show: it is still observational. A shrunken adjusted |
| 21 | +difference or a near_nogift group that resembles "away" would be evidence CONSISTENT with an energy/provisioning |
| 22 | +explanation; it would not exclude other explanations that also track energy or gift timing, and a surviving |
| 23 | +association after adjustment would show the explanation is incomplete, not that it is wrong. Gift receipt is |
| 24 | +measured from the change in every predator's energy around the gift call (InstrumentedMixin), so it is exact |
| 25 | +whatever the rules; energy bands are fixed edges, not fitted. |
| 26 | +
|
| 27 | +Example: |
| 28 | + python -m predpreygrass.non_evolutionary.predator_sexual_reproduction.analyze_provisioning \ |
| 29 | + --run CONTROL_s42=~/simulation_results/ray_results/PPO_FIXED_PREY_DENSITY_CONTROL_SEED42 --episodes 30 |
| 30 | +""" |
| 31 | +import argparse |
| 32 | +import glob |
| 33 | +import json |
| 34 | +import os |
| 35 | + |
| 36 | +import numpy as np |
| 37 | +import torch |
| 38 | + |
| 39 | +from predpreygrass.non_evolutionary.predator_sexual_reproduction.analyze_energy_sources import make_instrumented_env |
| 40 | +from predpreygrass.non_evolutionary.predator_sexual_reproduction.analyze_mate_contingency import mate_bucket |
| 41 | +from predpreygrass.non_evolutionary.predator_sexual_reproduction.analyze_prey_approach_from_checkpoint import ( |
| 42 | + POLICY_IDS, |
| 43 | + action_geometry, |
| 44 | + load_modules, |
| 45 | + module_probs, |
| 46 | + new_stats, |
| 47 | + policy_of, |
| 48 | + record, |
| 49 | +) |
| 50 | + |
| 51 | +TARGETS = ("prey", "fruit") |
| 52 | +GROUPS = ("gift_recent", "near_nogift", "far", "dead", "abandoned") # virgin females are excluded |
| 53 | +ENERGY_EDGES = (4.0, 8.0, 12.0) # bands: <4, 4-8, 8-12, >=12 (predator_creation_energy_threshold is 12) |
| 54 | +N_BANDS = len(ENERGY_EDGES) + 1 |
| 55 | + |
| 56 | + |
| 57 | +def energy_band(e): |
| 58 | + return int(np.searchsorted(ENERGY_EDGES, e, side="right")) |
| 59 | + |
| 60 | + |
| 61 | +def run_episodes(env_config, modules, n_episodes, seed0, near_radius, gift_window): |
| 62 | + env = make_instrumented_env(env_config) |
| 63 | + moves = np.array([env.action_to_move_tuple[a] for a in range(env.num_actions)]) |
| 64 | + offset = (env.predator_obs_range - 1) // 2 |
| 65 | + sample_rng = np.random.default_rng(seed0) |
| 66 | + per_episode = [] |
| 67 | + for ep in range(n_episodes): |
| 68 | + observations, _ = env.reset(seed=seed0 + ep) |
| 69 | + active = list(observations.keys()) |
| 70 | + stats = {(t, g, b): new_stats() for t in TARGETS for g in GROUPS for b in range(N_BANDS)} |
| 71 | + energy_sum = {g: [0.0, 0] for g in GROUPS} |
| 72 | + last_gift_step = {} |
| 73 | + seen_gifts = {} |
| 74 | + step = 0 |
| 75 | + while True: |
| 76 | + targets = { |
| 77 | + "prey": np.array(list(env.prey_positions.values()), dtype=int).reshape(-1, 2), |
| 78 | + "fruit": np.array( |
| 79 | + [p for f, p in env.fruit_positions.items() if env.fruit_energies[f] > 0], dtype=int |
| 80 | + ).reshape(-1, 2), |
| 81 | + } |
| 82 | + actions = {} |
| 83 | + for pid in POLICY_IDS: |
| 84 | + agents = [a for a in active if policy_of(a) == pid] |
| 85 | + if not agents: |
| 86 | + continue |
| 87 | + rows = module_probs(modules[pid], [observations[a] for a in agents]) |
| 88 | + for i, agent in enumerate(agents): |
| 89 | + probs = rows[i] |
| 90 | + if pid == "predator_female_policy": |
| 91 | + bucket = mate_bucket(env, agent, near_radius) |
| 92 | + if bucket == "near": |
| 93 | + since = step - last_gift_step.get(agent, -10**9) |
| 94 | + group = "gift_recent" if since <= gift_window else "near_nogift" |
| 95 | + elif bucket in ("far", "dead", "abandoned"): |
| 96 | + group = bucket |
| 97 | + else: |
| 98 | + group = None # virgin |
| 99 | + if group is not None: |
| 100 | + energy = env.agent_energies[agent] |
| 101 | + band = energy_band(energy) |
| 102 | + energy_sum[group][0] += energy |
| 103 | + energy_sum[group][1] += 1 |
| 104 | + pos = np.array(env.agent_positions[agent], dtype=int) |
| 105 | + for t in TARGETS: |
| 106 | + geo = action_geometry(pos, targets[t], moves, env.grid_size, offset) |
| 107 | + if geo is not None: |
| 108 | + record(stats[(t, group, band)], geo, probs) |
| 109 | + actions[agent] = int(sample_rng.choice(env.num_actions, p=probs)) |
| 110 | + observations, _, terminations, truncations, _ = env.step(actions) |
| 111 | + step += 1 |
| 112 | + # Gift receipts observed during THIS step become visible to the NEXT decision. |
| 113 | + for a, total in list(env.gift_received.items()): |
| 114 | + if total > seen_gifts.get(a, 0.0) + 1e-12: |
| 115 | + last_gift_step[a] = step |
| 116 | + seen_gifts[a] = total |
| 117 | + active = [a for a in observations if not terminations.get(a, False) and not truncations.get(a, False)] |
| 118 | + if terminations.get("__all__") or truncations.get("__all__"): |
| 119 | + break |
| 120 | + per_episode.append((stats, energy_sum)) |
| 121 | + return per_episode |
| 122 | + |
| 123 | + |
| 124 | +def _arrays(per_episode): |
| 125 | + """num/den arrays of shape (episodes, groups, bands) per target.""" |
| 126 | + E = len(per_episode) |
| 127 | + out = {} |
| 128 | + for t in TARGETS: |
| 129 | + num = np.zeros((E, len(GROUPS), N_BANDS)) |
| 130 | + den = np.zeros((E, len(GROUPS), N_BANDS)) |
| 131 | + for e, (stats, _) in enumerate(per_episode): |
| 132 | + for gi, g in enumerate(GROUPS): |
| 133 | + for b in range(N_BANDS): |
| 134 | + s = stats[(t, g, b)] |
| 135 | + num[e, gi, b], den[e, gi, b] = s["bias"], s["n"] |
| 136 | + out[t] = (num, den) |
| 137 | + return out |
| 138 | + |
| 139 | + |
| 140 | +GI = {g: i for i, g in enumerate(GROUPS)} |
| 141 | +NEAR = [GI["gift_recent"], GI["near_nogift"]] |
| 142 | +AWAY = [GI["far"], GI["dead"], GI["abandoned"]] |
| 143 | + |
| 144 | + |
| 145 | +def _ratio(num, den, groups): |
| 146 | + n = num[..., groups, :].sum(axis=-2).sum(axis=-1) |
| 147 | + d = den[..., groups, :].sum(axis=-2).sum(axis=-1) |
| 148 | + with np.errstate(invalid="ignore", divide="ignore"): |
| 149 | + return np.where(d > 0, n / d, np.nan) |
| 150 | + |
| 151 | + |
| 152 | +def _adjusted(num, den): |
| 153 | + """Energy-band-standardized near-minus-away. Works on (..., groups, bands) totals.""" |
| 154 | + nn, nd = num[..., NEAR, :].sum(axis=-2), den[..., NEAR, :].sum(axis=-2) |
| 155 | + an, ad = num[..., AWAY, :].sum(axis=-2), den[..., AWAY, :].sum(axis=-2) |
| 156 | + with np.errstate(invalid="ignore", divide="ignore"): |
| 157 | + diff = nn / nd - an / ad |
| 158 | + valid = (nd > 0) & (ad > 0) |
| 159 | + w = np.where(valid, nd + ad, 0.0) |
| 160 | + diff = np.where(valid, diff, 0.0) |
| 161 | + tot = w.sum(axis=-1) |
| 162 | + with np.errstate(invalid="ignore", divide="ignore"): |
| 163 | + return np.where(tot > 0, (w * diff).sum(axis=-1) / tot, np.nan) |
| 164 | + |
| 165 | + |
| 166 | +def _stats_from_totals(num, den): |
| 167 | + """num/den shape (..., groups, bands) -> dict of named statistics, each shape (...).""" |
| 168 | + g = lambda name: [GI[name]] |
| 169 | + r = lambda groups: _ratio(num, den, groups) |
| 170 | + return { |
| 171 | + "crude near-away": r(NEAR) - r(AWAY), |
| 172 | + "energy-adjusted near-away": _adjusted(num, den), |
| 173 | + "near_nogift-away": r(g("near_nogift")) - r(AWAY), |
| 174 | + "gift_recent-away": r(g("gift_recent")) - r(AWAY), |
| 175 | + "gift_recent-near_nogift": r(g("gift_recent")) - r(g("near_nogift")), |
| 176 | + } |
| 177 | + |
| 178 | + |
| 179 | +def summarize(label, per_episode, rng, n_boot=2000): |
| 180 | + lines = [] |
| 181 | + E = len(per_episode) |
| 182 | + en = {g: (sum(ep[1][g][0] for ep in per_episode), sum(ep[1][g][1] for ep in per_episode)) for g in GROUPS} |
| 183 | + lines.append( |
| 184 | + f"{label}: female decisions by group (n, mean own energy): " |
| 185 | + + ", ".join(f"{g} n={n} e={(s / n if n else float('nan')):.1f}" for g, (s, n) in en.items()) |
| 186 | + ) |
| 187 | + arrs = _arrays(per_episode) |
| 188 | + for t in TARGETS: |
| 189 | + num, den = arrs[t] |
| 190 | + point = _stats_from_totals(num.sum(axis=0), den.sum(axis=0)) |
| 191 | + idx = rng.integers(0, E, size=(n_boot, E)) |
| 192 | + boot = _stats_from_totals(num[idx].sum(axis=1), den[idx].sum(axis=1)) |
| 193 | + parts = [] |
| 194 | + for name, val in point.items(): |
| 195 | + b = boot[name][~np.isnan(boot[name])] |
| 196 | + if len(b) >= 100 and not np.isnan(val): |
| 197 | + lo, hi = np.percentile(b, [2.5, 97.5]) |
| 198 | + parts.append(f"{name}={val:+.3f} [{lo:+.3f},{hi:+.3f}]") |
| 199 | + else: |
| 200 | + parts.append(f"{name}=n/a") |
| 201 | + lines.append(f" {t:5}: " + " | ".join(parts)) |
| 202 | + return lines |
| 203 | + |
| 204 | + |
| 205 | +def main(): |
| 206 | + parser = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter) |
| 207 | + parser.add_argument("--run", action="append", default=[], help="LABEL=path to a ray_results experiment dir") |
| 208 | + parser.add_argument("--checkpoint", type=int, default=29) |
| 209 | + parser.add_argument("--episodes", type=int, default=30) |
| 210 | + parser.add_argument("--seed", type=int, default=4000) |
| 211 | + parser.add_argument("--near-radius", type=int, default=3, help="also the gift range in this module's defaults") |
| 212 | + parser.add_argument("--gift-window", type=int, default=10, help="steps after a received gift counted as 'recent'") |
| 213 | + args = parser.parse_args() |
| 214 | + torch.set_num_threads(2) |
| 215 | + rng = np.random.default_rng(0) |
| 216 | + for spec in args.run: |
| 217 | + label, path = spec.split("=", 1) |
| 218 | + path = os.path.expanduser(path) |
| 219 | + config = json.load(open(os.path.join(path, "run_config.json")))["config_env"] |
| 220 | + trial = sorted(glob.glob(os.path.join(path, "PPO_*/")))[0] |
| 221 | + modules = load_modules(os.path.join(trial, f"checkpoint_{args.checkpoint:06d}")) |
| 222 | + per_episode = run_episodes(config, modules, args.episodes, args.seed, args.near_radius, args.gift_window) |
| 223 | + for line in summarize(f"{label} ckpt{args.checkpoint}", per_episode, rng): |
| 224 | + print(line, flush=True) |
| 225 | + print(flush=True) |
| 226 | + |
| 227 | + |
| 228 | +if __name__ == "__main__": |
| 229 | + main() |
0 commit comments