MCPcopy Create free account
hub / github.com/CompVis/zigma / token_drop

Method token_drop

model_zigma.py:292–303  ·  view source on GitHub ↗

Drops labels to enable classifier-free guidance.

(self, labels, force_drop_ids=None)

Source from the content-addressed store, hash-verified

290 self.dropout_prob = dropout_prob
291
292 def token_drop(self, labels, force_drop_ids=None):
293 """
294 Drops labels to enable classifier-free guidance.
295 """
296 if force_drop_ids is None:
297 drop_ids = (
298 torch.rand(labels.shape[0], device=labels.device) < self.dropout_prob
299 )
300 else:
301 drop_ids = force_drop_ids == 1
302 labels = torch.where(drop_ids, self.num_classes, labels)
303 return labels
304
305 def forward(self, labels, train, force_drop_ids=None):
306 use_dropout = self.dropout_prob > 0

Callers 1

forwardMethod · 0.95

Calls

no outgoing calls

Tested by

no test coverage detected