MCPcopy Create free account
hub / github.com/NineAbyss/ZeroG / DataWrapper

Class DataWrapper

code/dataset_benchmark.py:26–64  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

24
25
26class 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

Callers 1

load_node_datasetFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected