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

Method add_batch

lib/layerwrapper.py:25–49  ·  view source on GitHub ↗

tar: batch_size * seq_len, inp corresponding to the position where tar == -100 will be ignored

(self, inp, out, tar)

Source from the content-addressed store, hash-verified

23 self.layer_name = layer_name
24
25 def add_batch(self, inp, out, tar):
26 """
27 tar: batch_size * seq_len, inp corresponding to the position where tar == -100 will be ignored
28 """
29 if len(inp.shape) == 2:
30 inp = inp.unsqueeze(0)
31 if len(tar.shape) == 2:
32 tar = tar.unsqueeze(0)
33
34 tmp = inp.shape[0] # bs
35
36 mask = tar.ne(-100)
37 if isinstance(self.layer, nn.Linear):
38 if len(inp.shape) == 3:
39 inp = inp.reshape((-1, inp.shape[-1]))
40 mask = mask.flatten()
41 inp = inp[mask] # remove -100's
42 inp = inp.t()
43
44 self.scaler_row *= self.nsamples / (self.nsamples + tmp)
45 self.nsamples += tmp
46
47 inp = inp.type(torch.float32)
48 self.scaler_row += torch.norm(inp, p=2, dim=1) ** 2 / self.nsamples
49 self.activations.append(inp)

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected