| 24 | |
| 25 | |
| 26 | class DataWrapper: |
| 27 | def __init__(self, data, args): |
| 28 | self._data = data |
| 29 | self.args = args |
| 30 | # self.emb_file = 'datasets/wikics/processed/wikics_text_embedding.pt' |
| 31 | # self.label_embedding() |
| 32 | # if args.if_text and args.cl: |
| 33 | # self.feature_embedding() |
| 34 | # self.label_embedding() |
| 35 | self.x = self._data.x |
| 36 | # self.label_emb = self._data.label_emb |
| 37 | self.raw_texts = self._data.raw_texts |
| 38 | self.label_text = self._data.label_name |
| 39 | |
| 40 | @property |
| 41 | def data(self): |
| 42 | return self._data |
| 43 | |
| 44 | def label_embedding(self): |
| 45 | text_model = TextModel(self.args.text_encoder) |
| 46 | text_features = [] |
| 47 | raw_texts = self.data.label_name |
| 48 | for text in tqdm.tqdm(raw_texts, desc="Processing label texts"): |
| 49 | text_features.append(text_model(text).unsqueeze(dim=0).cpu()) |
| 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 |