| 90 | |
| 91 | |
| 92 | class TextGCNTrainer: |
| 93 | def __init__(self, args, model, pre_data): |
| 94 | self.args = args |
| 95 | self.model = model |
| 96 | self.device = args.device |
| 97 | |
| 98 | self.max_epoch = self.args.max_epoch |
| 99 | self.set_seed() |
| 100 | |
| 101 | self.dataset = args.dataset |
| 102 | self.predata = pre_data |
| 103 | self.earlystopping = EarlyStopping(args.early_stopping) |
| 104 | |
| 105 | def set_seed(self): |
| 106 | th.manual_seed(self.args.seed) |
| 107 | np.random.seed(self.args.seed) |
| 108 | |
| 109 | def fit(self): |
| 110 | self.prepare_data() |
| 111 | self.model = self.model(nfeat=self.nfeat_dim, |
| 112 | nhid=self.args.nhid, |
| 113 | nclass=self.nclass, |
| 114 | dropout=self.args.dropout) |
| 115 | print(self.model.parameters) |
| 116 | self.model = self.model.to(self.device) |
| 117 | |
| 118 | self.optimizer = th.optim.Adam(self.model.parameters(), lr=self.args.lr) |
| 119 | self.criterion = th.nn.CrossEntropyLoss() |
| 120 | |
| 121 | self.model_param = sum(param.numel() for param in self.model.parameters()) |
| 122 | print('# model parameters:', self.model_param) |
| 123 | self.convert_tensor() |
| 124 | |
| 125 | start = time() |
| 126 | self.train() |
| 127 | self.train_time = time() - start |
| 128 | |
| 129 | @classmethod |
| 130 | def set_description(cls, desc): |
| 131 | string = "" |
| 132 | for key, value in desc.items(): |
| 133 | if isinstance(value, int): |
| 134 | string += f"{key}:{value} " |
| 135 | else: |
| 136 | string += f"{key}:{value:.4f} " |
| 137 | print(string) |
| 138 | |
| 139 | def prepare_data(self): |
| 140 | self.adj = self.predata.adj |
| 141 | self.nfeat_dim = self.predata.nfeat_dim |
| 142 | self.features = self.predata.features |
| 143 | self.target = self.predata.target |
| 144 | self.nclass = self.predata.nclass |
| 145 | |
| 146 | self.train_lst, self.val_lst = train_test_split(self.predata.train_lst, |
| 147 | test_size=self.args.val_ratio, |
| 148 | shuffle=True, |
| 149 | random_state=self.args.seed) |