| 50 | self.data.label_emb = torch.cat(text_features, dim=0) |
| 51 | |
| 52 | def feature_embedding(self): |
| 53 | emb_file = f"saved_embs/{self.data.x.shape[0]}.pt" |
| 54 | if not os.path.exists(emb_file): |
| 55 | text_model = TextModel(self.args.text_encoder) |
| 56 | text_features = [] |
| 57 | raw_texts = self.data.raw_texts |
| 58 | |
| 59 | for text in tqdm.tqdm(raw_texts, desc="Processing node texts"): |
| 60 | text_features.append(text_model(text).unsqueeze(dim=0).cpu()) |
| 61 | self.data.x = torch.cat(text_features, dim=0) |
| 62 | torch.save(self.data.x, emb_file) |
| 63 | else: |
| 64 | self.data.x = torch.load(emb_file) |
| 65 | |
| 66 | |
| 67 | # ======================================================================f |