You can not select more than 25 topics Topics must start with a chinese character,a letter or number, can include dashes ('-') and can be up to 35 characters long.

node_classification_full.py 16 kB

5 years ago
5 years ago
5 years ago
5 years ago
5 years ago
5 years ago
5 years ago
5 years ago
5 years ago
5 years ago
5 years ago
5 years ago
5 years ago
5 years ago
5 years ago
5 years ago
5 years ago
5 years ago
5 years ago
5 years ago
5 years ago
123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521
  1. """
  2. Node classification Full Trainer Implementation
  3. """
  4. from . import register_trainer
  5. from .base import BaseNodeClassificationTrainer, EarlyStopping, Evaluation
  6. import torch
  7. from torch.optim.lr_scheduler import (
  8. StepLR,
  9. MultiStepLR,
  10. ExponentialLR,
  11. ReduceLROnPlateau,
  12. )
  13. import torch.nn.functional as F
  14. from ..model import MODEL_DICT, BaseModel
  15. from .evaluation import get_feval, Logloss
  16. from typing import Union
  17. from copy import deepcopy
  18. from ...utils import get_logger
  19. LOGGER = get_logger("node classification trainer")
  20. @register_trainer("NodeClassificationFull")
  21. class NodeClassificationFullTrainer(BaseNodeClassificationTrainer):
  22. """
  23. The node classification trainer.
  24. Used to automatically train the node classification problem.
  25. Parameters
  26. ----------
  27. model: ``BaseModel`` or ``str``
  28. The (name of) model used to train and predict.
  29. optimizer: ``Optimizer`` of ``str``
  30. The (name of) optimizer used to train and predict.
  31. lr: ``float``
  32. The learning rate of node classification task.
  33. max_epoch: ``int``
  34. The max number of epochs in training.
  35. early_stopping_round: ``int``
  36. The round of early stop.
  37. device: ``torch.device`` or ``str``
  38. The device where model will be running on.
  39. init: ``bool``
  40. If True(False), the model will (not) be initialized.
  41. """
  42. def __init__(
  43. self,
  44. model: Union[BaseModel, str]=None,
  45. num_features=None,
  46. num_classes=None,
  47. optimizer=None,
  48. lr=None,
  49. max_epoch=None,
  50. early_stopping_round=None,
  51. weight_decay=1e-4,
  52. device="auto",
  53. init=True,
  54. feval=[Logloss],
  55. loss="nll_loss",
  56. lr_scheduler_type=None,
  57. *args,
  58. **kwargs
  59. ):
  60. super().__init__(
  61. model,
  62. num_features,
  63. num_classes,
  64. device=device,
  65. init=init,
  66. feval=feval,
  67. loss=loss,
  68. )
  69. self.opt_received = optimizer
  70. if type(optimizer) == str and optimizer.lower() == "adam":
  71. self.optimizer = torch.optim.Adam
  72. elif type(optimizer) == str and optimizer.lower() == "sgd":
  73. self.optimizer = torch.optim.SGD
  74. else:
  75. self.optimizer = torch.optim.Adam
  76. self.lr_scheduler_type = lr_scheduler_type
  77. self.lr = lr if lr is not None else 1e-4
  78. self.max_epoch = max_epoch if max_epoch is not None else 100
  79. self.early_stopping_round = (
  80. early_stopping_round if early_stopping_round is not None else 100
  81. )
  82. self.args = args
  83. self.kwargs = kwargs
  84. self.feval = get_feval(feval)
  85. self.weight_decay = weight_decay
  86. self.early_stopping = EarlyStopping(
  87. patience=early_stopping_round, verbose=False
  88. )
  89. self.valid_result = None
  90. self.valid_result_prob = None
  91. self.valid_score = None
  92. self.initialized = False
  93. self.space = [
  94. {
  95. "parameterName": "max_epoch",
  96. "type": "INTEGER",
  97. "maxValue": 500,
  98. "minValue": 10,
  99. "scalingType": "LINEAR",
  100. },
  101. {
  102. "parameterName": "early_stopping_round",
  103. "type": "INTEGER",
  104. "maxValue": 30,
  105. "minValue": 10,
  106. "scalingType": "LINEAR",
  107. },
  108. {
  109. "parameterName": "lr",
  110. "type": "DOUBLE",
  111. "maxValue": 1e-1,
  112. "minValue": 1e-4,
  113. "scalingType": "LOG",
  114. },
  115. {
  116. "parameterName": "weight_decay",
  117. "type": "DOUBLE",
  118. "maxValue": 1e-2,
  119. "minValue": 1e-4,
  120. "scalingType": "LOG",
  121. },
  122. ]
  123. self.hyperparams = {
  124. "max_epoch": self.max_epoch,
  125. "early_stopping_round": self.early_stopping_round,
  126. "lr": self.lr,
  127. "weight_decay": self.weight_decay,
  128. }
  129. if init is True:
  130. self.initialize()
  131. def initialize(self):
  132. # Initialize the auto model in trainer.
  133. if self.initialized is True:
  134. return
  135. self.initialized = True
  136. self.model.initialize()
  137. def get_model(self):
  138. # Get auto model used in trainer.
  139. return self.model
  140. @classmethod
  141. def get_task_name(cls):
  142. # Get task name, i.e., `NodeClassification`.
  143. return "NodeClassification"
  144. def train_only(self, data, train_mask=None):
  145. """
  146. The function of training on the given dataset and mask.
  147. Parameters
  148. ----------
  149. data: The node classification dataset used to be trained. It should consist of masks, including train_mask, and etc.
  150. train_mask: The mask used in training stage.
  151. Returns
  152. -------
  153. self: ``autogl.train.NodeClassificationTrainer``
  154. A reference of current trainer.
  155. """
  156. data = data.to(self.device)
  157. mask = data.train_mask if train_mask is None else train_mask
  158. optimizer = self.optimizer(
  159. self.model.parameters(), lr=self.lr, weight_decay=self.weight_decay
  160. )
  161. # scheduler = StepLR(optimizer, step_size=100, gamma=0.1)
  162. lr_scheduler_type = self.lr_scheduler_type
  163. if type(lr_scheduler_type) == str and lr_scheduler_type == "steplr":
  164. scheduler = StepLR(optimizer, step_size=100, gamma=0.1)
  165. elif type(lr_scheduler_type) == str and lr_scheduler_type == "multisteplr":
  166. scheduler = MultiStepLR(optimizer, milestones=[30, 80], gamma=0.1)
  167. elif type(lr_scheduler_type) == str and lr_scheduler_type == "exponentiallr":
  168. scheduler = ExponentialLR(optimizer, gamma=0.1)
  169. elif (
  170. type(lr_scheduler_type) == str and lr_scheduler_type == "reducelronplateau"
  171. ):
  172. scheduler = ReduceLROnPlateau(optimizer, "min")
  173. else:
  174. scheduler = None
  175. for epoch in range(1, self.max_epoch):
  176. self.model.model.train()
  177. optimizer.zero_grad()
  178. res = self.model.model.forward(data)
  179. if hasattr(F, self.loss):
  180. loss = getattr(F, self.loss)(res[mask], data.y[mask])
  181. else:
  182. raise TypeError(
  183. "PyTorch does not support loss type {}".format(self.loss)
  184. )
  185. loss.backward()
  186. optimizer.step()
  187. if self.lr_scheduler_type:
  188. scheduler.step()
  189. if hasattr(data, "val_mask") and data.val_mask is not None:
  190. if type(self.feval) is list:
  191. feval = self.feval[0]
  192. else:
  193. feval = self.feval
  194. val_loss = self.evaluate([data], mask=data.val_mask, feval=feval)
  195. if feval.is_higher_better() is True:
  196. val_loss = -val_loss
  197. self.early_stopping(val_loss, self.model.model)
  198. if self.early_stopping.early_stop:
  199. LOGGER.debug("Early stopping at %d", epoch)
  200. break
  201. if hasattr(data, "val_mask") and data.val_mask is not None:
  202. self.early_stopping.load_checkpoint(self.model.model)
  203. def predict_only(self, data, test_mask=None):
  204. """
  205. The function of predicting on the given dataset and mask.
  206. Parameters
  207. ----------
  208. data: The node classification dataset used to be predicted.
  209. train_mask: The mask used in training stage.
  210. Returns
  211. -------
  212. res: The result of predicting on the given dataset.
  213. """
  214. # mask = data.test_mask if test_mask is None else test_mask
  215. data = data.to(self.device)
  216. self.model.model.eval()
  217. with torch.no_grad():
  218. res = self.model.model.forward(data)
  219. return res
  220. def train(self, dataset, keep_valid_result=True):
  221. """
  222. The function of training on the given dataset and keeping valid result.
  223. Parameters
  224. ----------
  225. dataset: The node classification dataset used to be trained.
  226. keep_valid_result: ``bool``
  227. If True(False), save the validation result after training.
  228. Returns
  229. -------
  230. self: ``autogl.train.NodeClassificationTrainer``
  231. A reference of current trainer.
  232. """
  233. data = dataset[0]
  234. self.train_only(data)
  235. if keep_valid_result:
  236. self.valid_result = self.predict_only(data)[data.val_mask].max(1)[1]
  237. self.valid_result_prob = self.predict_only(data)[data.val_mask]
  238. self.valid_score = self.evaluate(
  239. dataset, mask=data.val_mask, feval=self.feval
  240. )
  241. def predict(self, dataset, mask=None):
  242. """
  243. The function of predicting on the given dataset.
  244. Parameters
  245. ----------
  246. dataset: The node classification dataset used to be predicted.
  247. mask: ``train``, ``val``, or ``test``.
  248. The dataset mask.
  249. Returns
  250. -------
  251. The prediction result of ``predict_proba``.
  252. """
  253. return self.predict_proba(dataset, mask=mask, in_log_format=True).max(1)[1]
  254. def predict_proba(self, dataset, mask=None, in_log_format=False):
  255. """
  256. The function of predicting the probability on the given dataset.
  257. Parameters
  258. ----------
  259. dataset: The node classification dataset used to be predicted.
  260. mask: ``train``, ``val``, or ``test``.
  261. The dataset mask.
  262. in_log_format: ``bool``.
  263. If True(False), the probability will (not) be log format.
  264. Returns
  265. -------
  266. The prediction result.
  267. """
  268. data = dataset[0]
  269. data = data.to(self.device)
  270. if mask is not None:
  271. if mask == "val":
  272. mask = data.val_mask
  273. elif mask == "test":
  274. mask = data.test_mask
  275. elif mask == "train":
  276. mask = data.train_mask
  277. else:
  278. mask = data.test_mask
  279. ret = self.predict_only(data, mask)[mask]
  280. if in_log_format is True:
  281. return ret
  282. else:
  283. return torch.exp(ret)
  284. def get_valid_predict(self):
  285. # """Get the valid result."""
  286. return self.valid_result
  287. def get_valid_predict_proba(self):
  288. # """Get the valid result (prediction probability)."""
  289. return self.valid_result_prob
  290. def get_valid_score(self, return_major=True):
  291. """
  292. The function of getting the valid score.
  293. Parameters
  294. ----------
  295. return_major: ``bool``.
  296. If True, the return only consists of the major result.
  297. If False, the return consists of the all results.
  298. Returns
  299. -------
  300. result: The valid score in training stage.
  301. """
  302. if isinstance(self.feval, list):
  303. if return_major:
  304. return self.valid_score[0], self.feval[0].is_higher_better()
  305. else:
  306. return self.valid_score, [f.is_higher_better() for f in self.feval]
  307. else:
  308. return self.valid_score, self.feval.is_higher_better()
  309. def get_name_with_hp(self):
  310. # """Get the name of hyperparameter."""
  311. name = "-".join(
  312. [
  313. str(self.optimizer),
  314. str(self.lr),
  315. str(self.max_epoch),
  316. str(self.early_stopping_round),
  317. str(self.model),
  318. str(self.device),
  319. ]
  320. )
  321. name = (
  322. name
  323. + "|"
  324. + "-".join(
  325. [
  326. str(x[0]) + "-" + str(x[1])
  327. for x in self.model.get_hyper_parameter().items()
  328. ]
  329. )
  330. )
  331. return name
  332. def evaluate(self, dataset, mask=None, feval=None):
  333. """
  334. The function of training on the given dataset and keeping valid result.
  335. Parameters
  336. ----------
  337. dataset: The node classification dataset used to be evaluated.
  338. mask: ``train``, ``val``, or ``test``.
  339. The dataset mask.
  340. feval: ``str``.
  341. The evaluation method used in this function.
  342. Returns
  343. -------
  344. res: The evaluation result on the given dataset.
  345. """
  346. data = dataset[0]
  347. data = data.to(self.device)
  348. test_mask = mask
  349. if feval is None:
  350. feval = self.feval
  351. else:
  352. feval = get_feval(feval)
  353. if test_mask is None:
  354. test_mask = data.test_mask
  355. elif test_mask == "test":
  356. test_mask = data.test_mask
  357. elif test_mask == "val":
  358. test_mask = data.val_mask
  359. elif test_mask == "train":
  360. test_mask = data.train_mask
  361. y_pred_prob = self.predict_proba(dataset, mask)
  362. y_pred = y_pred_prob.max(1)[1]
  363. y_true = data.y[test_mask]
  364. if not isinstance(feval, list):
  365. feval = [feval]
  366. return_signle = True
  367. else:
  368. return_signle = False
  369. res = []
  370. for f in feval:
  371. try:
  372. res.append(f.evaluate(y_pred_prob, y_true))
  373. except:
  374. res.append(f.evaluate(y_pred_prob.cpu().numpy(), y_true.cpu().numpy()))
  375. if return_signle:
  376. return res[0]
  377. return res
  378. def to(self, new_device):
  379. assert isinstance(new_device, torch.device)
  380. self.device = new_device
  381. if self.model is not None:
  382. self.model.to(self.device)
  383. def duplicate_from_hyper_parameter(self, hp: dict, model=None, restricted=True):
  384. """
  385. The function of duplicating a new instance from the given hyperparameter.
  386. Parameters
  387. ----------
  388. hp: ``dict``.
  389. The hyperparameter used in the new instance.
  390. model: The model used in the new instance of trainer.
  391. restricted: ``bool``.
  392. If False(True), the hyperparameter should (not) be updated from origin hyperparameter.
  393. Returns
  394. -------
  395. self: ``autogl.train.NodeClassificationTrainer``
  396. A new instance of trainer.
  397. """
  398. if not restricted:
  399. origin_hp = deepcopy(self.hyperparams)
  400. origin_hp.update(hp)
  401. hp = origin_hp
  402. if model is None:
  403. model = self.model
  404. model = model.from_hyper_parameter(
  405. dict(
  406. [
  407. x
  408. for x in hp.items()
  409. if x[0] in [y["parameterName"] for y in model.space]
  410. ]
  411. )
  412. )
  413. ret = self.__class__(
  414. model=model,
  415. num_features=self.num_features,
  416. num_classes=self.num_classes,
  417. optimizer=self.opt_received,
  418. lr=hp["lr"],
  419. max_epoch=hp["max_epoch"],
  420. early_stopping_round=hp["early_stopping_round"],
  421. device=self.device,
  422. weight_decay=hp["weight_decay"],
  423. feval=self.feval,
  424. loss=self.loss,
  425. lr_scheduler_type=self.lr_scheduler_type,
  426. init=True,
  427. *self.args,
  428. **self.kwargs
  429. )
  430. return ret
  431. @property
  432. def hyper_parameter_space(self):
  433. # """Get the space of hyperparameter."""
  434. return self.space
  435. @hyper_parameter_space.setter
  436. def hyper_parameter_space(self, space):
  437. # """Set the space of hyperparameter."""
  438. self.space = space
  439. def get_hyper_parameter(self):
  440. # """Get the hyperparameter in this trainer."""
  441. return self.hyperparams