Repeated module calls belong to TDHook#
A TDHook Target selects a zero-based occurrence directly. XDRL does not install internal hooks or maintain call counters.
[1]:
import torch
from tensordict import TensorDict
from tensordict.nn import TensorDictModule
from tdhook.latent import ActivationCaching
from tdhook.targets import Target
from tdhook.workflow import Workflow
from xdrl import interpret
class ReusedLayer(torch.nn.Module):
def __init__(self):
super().__init__()
self.shared = torch.nn.Identity()
def forward(self, value):
return self.shared(value + 1) + self.shared(value + 2)
policy = TensorDictModule(ReusedLayer(), in_keys=["observation"], out_keys=["action"])
component = interpret(policy)
target = Target("module.shared", "activation", -1, (0,), occurrences=(1,))
workflow = Workflow(ActivationCaching(target, cache_key=("activations", "selected")))
data = TensorDict({"observation": torch.tensor([[1.0, 2.0]])}, batch_size=[1])
result = component.run(workflow, data)
torch.testing.assert_close(result.data["activations", "selected", "module.shared"], torch.tensor([[3.0]]))