Compose an end-to-end policy investigation#
This bounded, deterministic, CPU-only tutorial follows one frozen TorchRL policy through a complete evidence chain:
evaluate competence on every non-terminal state in a corridor;
identify a hidden channel associated with the goal direction;
trace that channel’s relevance back to the observation;
intervene on the selected channel and matched channel controls;
measure open-loop decisions and closed-loop behavior.
The example adapts the useful corridor fixture from TDHook PR #105. TensorDict and TorchRL own the policy boundary, TDHook owns the model-internal methods, and XDRL supplies the single bridge between them. The result applies only to the synthetic fitted-Q policy trained below; it is not a paper reproduction or evidence about navigation policies in general.
[1]:
import torch
from tensordict import TensorDict
from tensordict.nn import TensorDictModule
from torch import nn
from tdhook.attribution import LRP
from tdhook.concepts import ChannelConditionedLRP, ConceptSelection
from tdhook.latent import ActivationCaching, SteeringVectors
from tdhook.targets import Target
from tdhook.workflow import Workflow
from xdrl import interpret
torch.manual_seed(7)
torch.set_num_threads(1)
CORRIDOR_LENGTH = 9
FEATURES = 8
SEED = 7
DEVICE = torch.device("cpu")
1. Freeze a complete decision panel#
An agent and goal occupy a nine-cell corridor. The two observation planes encode their positions; actions are 0 = left and 1 = right. Enumerating every non-terminal state and both transitions makes the evaluation set fixed rather than sampled.
[2]:
def encode_state(agent_position, goal_position):
observation = torch.zeros(2, CORRIDOR_LENGTH, device=DEVICE)
observation[0, agent_position] = 1.0
observation[1, goal_position] = 1.0
return observation
observations = []
next_observations = [[], []]
rewards = [[], []]
dones = [[], []]
optimal_actions = []
goal_is_right = []
state_pairs = []
for agent_position in range(CORRIDOR_LENGTH):
for goal_position in range(CORRIDOR_LENGTH):
if agent_position == goal_position:
continue
observations.append(encode_state(agent_position, goal_position))
optimal_actions.append(int(goal_position > agent_position))
goal_is_right.append(int(goal_position > agent_position))
state_pairs.append((agent_position, goal_position))
for action, displacement in ((0, -1), (1, 1)):
next_position = max(0, min(CORRIDOR_LENGTH - 1, agent_position + displacement))
done = next_position == goal_position
next_observations[action].append(encode_state(next_position, goal_position))
rewards[action].append(1.0 if done else -0.02)
dones[action].append(done)
observations = torch.stack(observations)
next_observations = torch.stack([torch.stack(values) for values in next_observations], dim=1)
rewards = torch.tensor(rewards, device=DEVICE).T
dones = torch.tensor(dones, device=DEVICE).T
optimal_actions = torch.tensor(optimal_actions, device=DEVICE)
goal_is_right = torch.tensor(goal_is_right, device=DEVICE)
evaluation_panel = TensorDict(
{"observation": observations, "target_action": optimal_actions, "concept_labels": goal_is_right},
batch_size=[len(observations)],
)
{"states": len(observations), "left_right_balance": torch.bincount(goal_is_right).tolist()}
[2]:
{'states': 72, 'left_right_balance': [36, 36]}
2. Fit, freeze, and expose the policy#
Fitted Q-iteration uses the complete transition table. After training, a public TorchRL TensorDictModule declares the input and output keys, and XDRL interpret exposes the same module to direct calls and TDHook workflows.
[3]:
class CorridorQNetwork(nn.Module):
def __init__(self):
super().__init__()
self.flatten = nn.Flatten()
self.features = nn.Sequential(
nn.Linear(2 * CORRIDOR_LENGTH, 16), nn.ReLU(), nn.Linear(16, FEATURES), nn.ReLU()
)
self.policy_head = nn.Linear(FEATURES, 2)
def forward(self, observation):
return self.policy_head(self.features(self.flatten(observation)))
q_network = CorridorQNetwork().to(DEVICE)
target_network = CorridorQNetwork().to(DEVICE)
target_network.load_state_dict(q_network.state_dict())
optimizer = torch.optim.Adam(q_network.parameters(), lr=0.02)
for iteration in range(600):
with torch.no_grad():
next_values = target_network(next_observations.flatten(0, 1)).max(-1).values.reshape(len(observations), 2)
bellman_targets = rewards + 0.95 * (~dones) * next_values
loss = nn.functional.smooth_l1_loss(q_network(observations), bellman_targets)
optimizer.zero_grad()
loss.backward()
optimizer.step()
if iteration % 10 == 0:
target_network.load_state_dict(q_network.state_dict())
q_network.eval()
class GreedyCorridorPolicy(nn.Module):
def __init__(self, q_network):
super().__init__()
self.q_network = q_network
def forward(self, observation):
action_value = self.q_network(observation)
return action_value, action_value.argmax(-1)
policy = TensorDictModule(
GreedyCorridorPolicy(q_network),
in_keys=["observation"],
out_keys=["action_value", "action"],
)
component = interpret(policy)
baseline_output = component(evaluation_panel.select("observation").clone())
baseline_actions = baseline_output["action"]
policy_accuracy = float((baseline_actions == optimal_actions).float().mean())
assert policy_accuracy == 1.0
{"complete_panel_accuracy": policy_accuracy}
[3]:
{'complete_panel_accuracy': 1.0}
[4]:
state_index = {pair: index for index, pair in enumerate(state_pairs)}
def rollout(action_table, start, goal, max_steps=12):
position = start
path = [position]
for _ in range(max_steps):
if position == goal:
break
action = int(action_table[state_index[(position, goal)]])
position = max(0, min(CORRIDOR_LENGTH - 1, position + (-1 if action == 0 else 1)))
path.append(position)
return path, position == goal
baseline_rollouts = [rollout(baseline_actions, start, goal) for start, goal in state_pairs]
baseline_success = sum(success for _path, success in baseline_rollouts) / len(baseline_rollouts)
assert baseline_success == 1.0
{
"closed_loop_success": baseline_success,
"left_goal_example": rollout(baseline_actions, 6, 1)[0],
"right_goal_example": rollout(baseline_actions, 2, 7)[0],
}
[4]:
{'closed_loop_success': 1.0,
'left_goal_example': [6, 5, 4, 3, 2, 1],
'right_goal_example': [2, 3, 4, 5, 6, 7]}
3. Run concept-conditioned attribution through XDRL#
The concept label says whether the goal lies to the right. LRP first attributes each chosen action value to the hidden layer. ConceptSelection then ranks channels by the contrast between right- and left-goal states, and channel-conditioned LRP traces the selected channel back to the input. This is an association and attribution result, not yet a causal claim.
XDRL validates the caller’s autograd state against TDHook’s public workflow plan and returns TDHook’s native result.
[5]:
def chosen_action_value(targets, additional):
values = targets["action_value"]
chosen = values.gather(-1, additional["target_action"].unsqueeze(-1)).squeeze(-1)
return TensorDict({"chosen_action_value": chosen}, batch_size=targets.batch_size)
lrp_options = {
"init_attr_targets": chosen_action_value,
"additional_init_keys": ["target_action"],
"warn_on_missing_rule": False,
}
diagnosis = Workflow(
LRP(
input_modules=["module.q_network.features.2"],
attribution_key=("attributions", "concept_examples"),
**lrp_options,
),
ConceptSelection(("attributions", "concept_examples", "module.q_network.features.2"), direction="negative"),
ChannelConditionedLRP(LRP(**lrp_options), condition_module="module.q_network.features.2"),
)
diagnostic_result = component.run(diagnosis, evaluation_panel.clone())
selection = diagnostic_result.data["metrics", "concept_selection"]
selected_channel = int(selection["channel"][0])
conditioned_relevance = diagnostic_result.data["attributions", "conditioned", "observation"]
assert selected_channel == 0
assert diagnostic_result.plan.model_passes == 2
{
"selected_left_associated_channel": selected_channel,
"selection_score": float(selection["score"][0]),
"conditioned_input_relevance_shape": tuple(conditioned_relevance.shape),
"model_passes": diagnostic_result.plan.model_passes,
}
[5]:
{'selected_left_associated_channel': 0,
'selection_score': -0.3201952278614044,
'conditioned_input_relevance_shape': (72, 2, 9),
'model_passes': 2}
4. Intervene with matched controls#
We first cache the complete hidden representation, then replace one channel with its mean activation in right-goal states. Every channel receives the same type of intervention and evaluation panel; non-selected channels are specificity controls.
The notebook owns the experimental comparison: both arms receive independent clones of the same frozen TensorDict and the same random seed. TDHook owns both intervention workflows, and XDRL runs each through the same interpreted component.
[6]:
capture = component.run(
Workflow(ActivationCaching("module.q_network.features.2", cache_key=("activations", "candidate_layer"))),
evaluation_panel.select("observation").clone(),
)
feature_values = capture.data["activations", "candidate_layer", "module.q_network.features.2"].detach()
right_context_means = feature_values[goal_is_right == 1].mean(0)
def no_op(*, output, **_):
return output
def replacement_callback(value):
def replace(*, output, **_):
return torch.full_like(output, value)
return replace
def run_channel_pair(channel):
target = Target("module.q_network.features.2", "activation", -1, (channel,))
baseline_workflow = Workflow(SteeringVectors([target], steer_fn=no_op))
intervention_workflow = Workflow(
SteeringVectors([target], steer_fn=replacement_callback(float(right_context_means[channel])))
)
inputs = evaluation_panel.select("observation")
torch.manual_seed(SEED)
baseline = component.run(baseline_workflow, inputs.clone())
torch.manual_seed(SEED)
intervention = component.run(intervention_workflow, inputs.clone())
return {"baseline": baseline, "intervention": intervention}
channel_pairs = {channel: run_channel_pair(channel) for channel in range(FEATURES)}
assert all(not child._forward_hooks for child in policy.modules())
{channel: pair["intervention"].plan.model_passes for channel, pair in channel_pairs.items()}
[6]:
{0: 1, 1: 1, 2: 1, 3: 1, 4: 1, 5: 1, 6: 1, 7: 1}
5. Evaluate decisions and behavior#
An activation change is not yet an RL result. We compare greedy actions over the complete panel, then use those frozen action tables in the corridor transition function. The selected-channel intervention should exceed every matched control and disrupt a left-goal rollout while preserving a right-goal rollout.
[7]:
pair_actions = {channel: pair["intervention"].data["action"] for channel, pair in channel_pairs.items()}
action_flips = {channel: int((actions != baseline_actions).sum()) for channel, actions in pair_actions.items()}
selected_actions = pair_actions[selected_channel]
selected_flips = action_flips[selected_channel]
max_matched_control_flips = max(flips for channel, flips in action_flips.items() if channel != selected_channel)
left_baseline = rollout(baseline_actions, 6, 1)
left_intervention = rollout(selected_actions, 6, 1)
right_baseline = rollout(baseline_actions, 2, 7)
right_intervention = rollout(selected_actions, 2, 7)
intervention_rollouts = [rollout(selected_actions, start, goal) for start, goal in state_pairs]
intervention_success = sum(success for _path, success in intervention_rollouts) / len(intervention_rollouts)
assert selected_flips > max_matched_control_flips
assert left_baseline[1] and not left_intervention[1]
assert right_baseline[1] and right_intervention[1]
behavioral_results = {
"action_flips_by_channel": action_flips,
"selected_channel": selected_channel,
"selected_channel_flips": selected_flips,
"largest_matched_control_flips": max_matched_control_flips,
"closed_loop_success": {"baseline": baseline_success, "intervention": intervention_success},
"left_goal": {"baseline": left_baseline, "intervention": left_intervention},
"right_goal": {"baseline": right_baseline, "intervention": right_intervention},
}
behavioral_results
[7]:
{'action_flips_by_channel': {0: 36, 1: 0, 2: 0, 3: 0, 4: 0, 5: 0, 6: 0, 7: 0},
'selected_channel': 0,
'selected_channel_flips': 36,
'largest_matched_control_flips': 0,
'closed_loop_success': {'baseline': 1.0, 'intervention': 0.5},
'left_goal': {'baseline': ([6, 5, 4, 3, 2, 1], True),
'intervention': ([6, 7, 8, 8, 8, 8, 8, 8, 8, 8, 8, 8, 8], False)},
'right_goal': {'baseline': ([2, 3, 4, 5, 6, 7], True),
'intervention': ([2, 3, 4, 5, 6, 7], True)}}
6. Results and scope#
Stage |
Evidence |
Supported interpretation |
|---|---|---|
Competence |
Complete-panel accuracy and all-pairs rollouts |
Competence in this finite corridor only |
Concept-conditioned attribution |
Selected channel and conditioned input relevance |
Association under the configured LRP rules |
Matched intervention |
Independent TensorDict clones, identical seeds, and explicit no-op baseline workflows |
Matched execution mechanics |
Specificity controls |
Identical interventions on every other hidden channel |
The selected intervention is more behaviorally specific in this policy |
Behavioral evaluation |
Complete-panel action changes and left/right rollouts |
A causal effect of this activation replacement on this frozen policy |
Nothing here supports a claim about a paper result, another checkpoint, another environment, or a general navigation mechanism. The reusable contribution is the explicit path from competence to a representational hypothesis, matched intervention, controls, and behavioral evaluation.