MCPcopy Create free account
hub / github.com/boyiwei/alignment-attribution-code / forward

Method forward

lib/model_wrapper.py:35–57  ·  view source on GitHub ↗
(self, x)

Source from the content-addressed store, hash-verified

33 self.n_samples = 0
34
35 def forward(self, x):
36 # TODO: normalize for numerical stability
37 # TODO: remove this after pruning
38
39 # DEBUG:
40 # print("input zero percentage", (x==0).sum() / x.numel() )
41
42 if self.record_activation:
43 if hasattr(self, "mask") and self.mask is not None:
44 x_ = x[self.mask]
45 else:
46 x_ = x
47
48 bs = x_.nelement() // x_.shape[-1]
49 self.activation_norms = self.activation_norms * (
50 self.n_samples / (self.n_samples + bs)
51 ) + (x_ * x_).view(-1, x_.shape[-1]).sum(dim=0) * (
52 1.0 / (self.n_samples + bs)
53 )
54 self.n_samples += bs
55
56 out = self.base(x)
57 return out
58
59
60class no_act_recording:

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected