(self, features, env=None)
| 713 | return {'loss': loss.item()} |
| 714 | |
| 715 | def update_embeddings_(self, features, env=None): |
| 716 | return_embedding = features.mean(0) |
| 717 | |
| 718 | if env is not None: |
| 719 | return_embedding = self.ema * return_embedding +\ |
| 720 | (1 - self.ema) * self.embeddings[env] |
| 721 | |
| 722 | self.embeddings[env] = return_embedding.clone().detach() |
| 723 | |
| 724 | return return_embedding.view(1, -1).repeat(len(features), 1) |
| 725 | |
| 726 | def predict(self, x, env=None): |
| 727 | features = self.featurizer(x) |