ãã®èšäºã¯ã NTTã³ãã¥ãã±ãŒã·ã§ã³ãº Advent Calendar 2022 20æ¥ç®ã®èšäºã§ãã ããã«ã¡ã¯ãã³ãã¥ãã±ãŒã·ã§ã³&ã¢ããªã±ãŒã·ã§ã³ãµãŒãã¹éšã®ç³äºã§ãã æ®æ®µã®æ¥åã§ã¯æç« èŠçŽæè¡ãçšããAPIãµãŒãã¹ 1 ã®éçºã»éçšã«åãçµãã§ãããŸãã ãã®èšäºã§ã¯ã°ã©ããã¥ãŒã©ã«ãããã¯ãŒã¯ïŒGNNïŒãç¹ã« Heterogeneous GraphïŒç°çš®ã°ã©ãïŒ ãæ±ã£ãGNNã«ã€ããŠç޹ä»ããŠããããšæããŸãã æ¬èšäºã§æ±ãå
容 ãã®èšäºã§åãæ±ãå
容ã¯ä»¥äžã§ãã ã°ã©ããã¥ãŒã©ã«ãããã¯ãŒã¯ïŒGNNïŒãšã¯ Heterogeneous GraphïŒç°çš®ã°ã©ãïŒ æ©æ¢°åŠç¿ã«ãããã°ã©ãããŒã¹ã®åé¡èšå® Pytorch-geometricã«ããã¢ãã«æ§ç¯ GNNã®æŠèŠãš Heterogeneous Graph ã«ã€ããŠç°¡åã«èª¬æãããåŸã«ãå®éã«ã¢ãã«ãäœæããŠããæµãã§å±éããŠãããŸãã æ¬èšäºã§ã¯ã¢ã«ãŽãªãºã ã®è©³çްãªè§£èª¬ãªã©ã¯çç¥ããŸãã®ã§ãããæ·±ãèå³ãããæ¹ã¯ãªã³ã¯ãã€ããŠãããŸãã®ã§è«æã解説èšäºãåç
§ããŠã¿ãŠãã ããã ã°ã©ããã¥ãŒã©ã«ãããã¯ãŒã¯ãšã¯ ã°ã©ããã¥ãŒã©ã«ãããã¯ãŒã¯ïŒGNNïŒãšã¯ã°ã©ãã§è¡šçŸãããããŒã¿ã深局åŠç¿ã§æ±ãããã®ãã¥ãŒã©ã«ãããã¯ãŒã¯ææ³ã®ç·ç§°ã§ããã°ã©ãããŒã¿ããè¡šçŸæœåºãããŠç®çã®ã¿ã¹ã¯ãè§£ããšããEnd2Endã¢ãããŒãã«ããæ©æ¢°åŠç¿ã¢ã«ãŽãªãºã ãšãªããŸããã¡ãžã£ãŒãªææ³ãšããŠã¯ GCN 2 ã GraphSAGE 3 ãªã©ããããŸãã GNN ã®ä»çµã¿ã«ã€ã㊠GCN ã®ææ³ãå
ã«ç°¡åã«èšèŒãããšãã°ã©ãã®é ç¹ïŒããŒãïŒã®ç¹åŸŽéã«å¯ŸããŠé£æ¥ããããŒãã®ç¹åŸŽéã«éã¿ãæãããã®ãå ããŠããæŒç®ãããããšã§ã察象ããŒãã«ã°ã©ãæ§é ã®æ
å ±ãå å³ãã衚çŸãç²åŸããããšãã£ãåããããŸãããã詳ãã説æã«ã€ããŠã¯ distill 4 ãšãããµã€ãã«ã Understanding Convolutions on Graphs ããšããGNNã®è§£èª¬èšäºãããã®ã§ãã¡ããåèã«ããŠã¿ãŠãã ããã Heterogeneous Graph Heterogeneous Graphãçè§£ããåæãšããŠãŸãã¯ã°ã©ãã®å®çŸ©ãã話ããŠãããŸãã ã°ã©ãçè«ã«ãããŠãã°ã©ããšã¯é ç¹ã瀺ãããŒããšãã®éã®é¢ä¿ã§ãããšããžãã衚çŸãããããŒã¿æ§é ã«ãªããŸããã€ãŸããããŒãéããšããžã§ç¹ãããšã«ãã£ãŠããŒãã«ãããããã¯ãŒã¯æ§é ã衚çŸãããã®ããäžè¬çã«ããç®ã«ããã°ã©ããšåŒã°ãããã®ã«ãªããŸãããããæ°åŒçã«å®çŸ©ãããšä»¥äžã®ããã«ãªããŸãã : ããŒãã®éåã : ãšããžã®éå ãããŠããã®ã°ã©ãã«ãããŠ1çš®é¡ã®ããŒããšãåãæå³åãã®ãšããžã«ãã£ãŠé¢ä¿ã瀺ãããã®ã Homogeneous Graph ãšåŒã³ãŸããäžæ¹ã§ãè€æ°ã®å€æ§ãªããŒããšãšããžãå«ãã§ããé¢ä¿ã®ã°ã©ãã Heterogeneous Graph ãšèšããŸããäŸãã°ããœãŒã·ã£ã«ãããã¯ãŒã¯ã®ãããªäººãšäººã®éã亀åé¢ä¿ã§ãªã³ã¯ããã°ã©ã㯠Homogeneous Graph ã§ãããåºèå©çšé¢ä¿ã®ãããªäººãšåºèã®éã賌買å®çžŸã§ãªã³ã¯ããã°ã©ã㯠Heterogeneous Graph ãšãªããŸããæ³èµ·ããããããã«å³ã§ç€ºããšä»¥äžã®ããã«ãªããŸãã ã¡ãªã¿ã« GNN ããŒã¹ã®ã¢ã«ãŽãªãºã ã®å€ãã¯å
¥åãåäžã®ããŒããšãšããžãæã€ãHomogeneous Graph ã察象ãšããææ³ãšãªã£ãŠããŸãïŒæè¿ã§ã¯ HAT 5 ã®ãã㪠Heterogeneous Graph ã察象ãšããã¢ã«ãŽãªãºã ãå¢ããŠããŠããŸãïŒããã®äžã§ããªã Heterogeneous Graph ãæ±ãå¿
èŠãããã®ããšããåãã«ã€ããŠã§ãããçŸå®äžçã«ãããŠèŠ³æž¬å¯Ÿè±¡ãã°ã©ã衚çŸã§æ§é åããããšããå Žåã«ãè€æ°ã®ããŒããŸãã¯ãšããžã«ããé¢ä¿ãå®çŸ©ããé »åºŠãé«ãããã§ãããã人ãšäººã®é¢ä¿ãæ§é åããå Žåãããããã人ãšç©ããµãŒãã¹ã®é¢ä¿ãæ§é åããå Žåã®æ¹ããã¿ãŒã³ãå€ãã®ã¯å®¹æã«æ³åã§ããŸãã ãšããããŸã§ Heterogeneous Graph ã®è©±ãããŠããŸãããããããã GNN èªäœã深局åŠç¿åéã®äžã§ãè¿å¹Žæ³šç®ãããŠããæè¡ã§ãããããŒã¿åæç«¶æåäŒã§ãã KDD CUP 2021 ã§ã¯ã OGB-LS ããšããã«ããŽãªã§3ã€ã®ã°ã©ãTaskãåãæ±ããããªã©ãã®æ³šç®åºŠã®é«ãã䌺ããŸãããŸããä»å¹Žéå¬ãããåœéäŒè°ã§ãã KDD 2022 ã®Research Trackã®äžã§ã¯å
š254ã®è«æã®äžã§çŽ80以äžãã®è«æãã°ã©ãã«é¢é£ããå
容ãåãæ±ããªã©ãã¬ã³ãé åã®1ã€ãšãªã£ãŠããããšãåãããŸããã ã°ã©ãã«ãããåé¡èšå® 次ã«ã°ã©ãã«ãããåé¡èšå®ã«ã€ããŠå°ãè§ŠããŸãã éåžžã®æ©æ¢°åŠç¿ã®åé¡èšå®ã§ã¯åé¡ãååž°ãšãã£ãã¿ã¹ã¯ãäžè¬çã«èããŸãããã°ã©ããæ±ãå Žåã«ã¯ã°ã©ãã«é©çšããåé¡èšå®ãæ€èšããå¿
èŠããããŸãããã®åé¡èšå®ã«ã¯å€§ãã3ã€ã®çš®é¡ããããŸãã 1ã€ã¯ããŒãã察象ãšããã¿ã¹ã¯ïŒNode CentricïŒã§ãã°ã©ãäžã«ãããããŒãåäœã®åé¡ãååž°ãšãã£ãã¿ã¹ã¯ãæ±ããŸãã2ã€ãã¯ã°ã©ãã察象ãšããã¿ã¹ã¯ïŒGraph CentricïŒã§ãã°ã©ãåäœãšããŠåé¡ãååž°ãšãã£ãã¿ã¹ã¯ãæ±ããŸããã°ã©ãåäœã®ã¿ã¹ã¯ã¯æŽ»çšã€ã¡ãŒãžããã¥ããããšæããŸãããååç©ã®åé¡ãªã©ã°ã©ããè€æ°ååšãããã¿ãŒã³ãæ³åããŠãããããšåãããããããšæããŸããæåŸã¯ãšããžã察象ãšããã¿ã¹ã¯ïŒEdge CentricïŒã§ãåã
ããŒãéã®ãšããžã«å¯ŸããŠäºæž¬ãããŠãããšããžã圢æãããã®ãããããšããžã®ã¯ã©ã¹ã¯äœãããšãã£ãã¿ã¹ã¯ãæ±ããŸãããã®ããã«ã°ã©ããæ±ãå Žåã«ã¯èªèº«ã®ç®çã«å¿ããã¿ã¹ã¯èšèšãå¿
èŠã«ãªããããåé¡èšå®ããã£ããæ€èšããäžã§å¿
èŠãªããŒã¿åéãå®è£
ãè¡ã£ãŠãããŸãã ãŸããããå°ãã¿ã¹ã¯ã®è£è¶³ãããŠãããšãäžèšã®åé¡èšå®ã«å ããŠãtrunsductiveããšãinductiveããšèšãåŠç¿ãšæšè«æã®ç¶æ³ã«ã€ããŠèæ
®ããŠããããšãéèŠã«ãªããŸãã ãtrunsductiveããšã¯åŠç¿ãšæšè«ã§åãã°ã©ããæ±ãå Žåã®ããšãæãããinductiveãã¯åŠç¿ããŒã¿ã«ãªãæ°ããã°ã©ããæ±ãå Žåã®ããšã瀺ããŸãïŒ semi-inductive 6 ãšãã£ãèãæ¹ãååšããŸãïŒããªããã®ãããªåé¡èšå®ã®éããæèããå¿
èŠãããããšãããšããã㯠GNN ã®ã¢ãã«æ§ç¯ã«ãŠéžæããã¢ã«ãŽãªãºã ãç°ãªã£ãŠããããã«ãªããŸã 7 ããã®ããããtrunsductiveããšãinductiveãã©ã¡ããã«ãã£ãŠéžæå¯èœãªã¢ã«ãŽãªãºã ã«å¶çŽãåºãŠããããšã«æ³šæããŠãã ããã ãã ãäžè¬çã«ã¯æ°ããæªç¥ã®ããŒãããšããžã«å¯ŸããŠäºæž¬ãè¡ããããšãã£ãå Žåã®æŽ»çšã·ãŒã³ã®æ¹ãå€ããšèããããããããinductiveããªåé¡èšå®ãããŒã¹ãšããŠèããŠããã°ãŸãã¯è¯ãããšæããŸãã å®éã«è©ŠããŠã¿ã ããããã¯å®éã« Heterogeneous Graph ãæ±ã£ã GNN ã®ã¢ãã«ãæ§ç¯ããŠã¿ãããšæããŸãã ä»å㯠Kaggle ã§å
¬éãããŠããã Recipes and Reviews ãã®ãªãŒãã³ããŒã¿ãå©çšããŸãããã®ããŒã¿ã¯ Food.com 8 ãšèšãæµ·å€ã®ã¬ã·ãå
±æãµã€ãããæçã¬ã·ããšãã®ã¬ã·ãã«å¯ŸããŠã®ãŠãŒã¶ã¬ãã¥ãŒã®æ
å ±ãããŒã¿åéãããã®ã«ãªããŸããæçã¬ã·ãã®ããŒã¿ã«ã¯ã¬ã·ãã«ãããã¡ã¿æ
å ±ãšå®éçãªæ é€çŽ ãšãã£ãç¹åŸŽéãå«ãã§ããããŠãŒã¶ã¬ãã¥ãŒã®ããŒã¿ã¯ãããŠãŒã¶ã該åœããæçã¬ã·ãã5段éã§è©äŸ¡ããå
容ãå«ãŸããŠããŸããããŒã¿ãµã€ãºã«ã€ããŠã500,000以äžã®ã¬ã·ãæ°ãš1,400,000ã®ã¬ãã¥ãŒæ°ãããããæ¯èŒçããªã¥ãŒã ã®å€§ããããŒã¿ãšãªã£ãŠããŸãã å©çšãããã¬ãŒã ã¯ãŒã¯ã§ããã Pytroch-geometric ãçšããŠå®è£
ãè¡ã£ãŠãããŸã 9 ãPytorch-geometricã§ã¯ãããŒãžã§ã³2.0.0ãã Heterogeneous Graph ããµããŒãããŠããŸãããã®ãããHeterogeneous Graph ã察象ãšããã¢ãã«ãäœæããå Žåã«ã¯ããŒãžã§ã³ã«æ³šæããŠãå©çšãã ããã ã§ã¯ãåé¡èšå®ãšããŠã¯ä»¥äžãèããŠã¿ãããšæããŸãã ãŠãŒã¶ãšã¬ã·ãã®é¢ä¿ã«ããäºéšã°ã©ãæ§é ã® Heterogeneous Graph ãå®çŸ©ããŠããŠãŒã¶ãã¬ã·ãã«èå³ã»é¢å¿ã瀺ãããäºæž¬ããã¿ã¹ã¯ïŒãªã³ã¯äºæž¬ïŒãè§£ãããšæããŸããããŒã¿å
容ã¯ãŠãŒã¶ã«ããã¬ã·ãã®è©äŸ¡ãæå³ãããããè©äŸ¡ãããšããäºå®ãèå³ã»é¢å¿ããããšããåºåãšããŠæ±ãã®ã¯å³å¯ã«ã¯æ£ããã¯ãªããšã¯æããŸãããä»åã¯äŸ¿å®äžãã®ãããªåé¡èšå®ãšããŸãã ããŒã¿æºå ã°ã©ãããŒã¿æŽåœ¢ ããŒã¿ã®èªã¿èŸŒã¿ããããŒãã«ããŒã¿ãã°ã©ãããŒã¿ã«å€æããããã®åŠçãèšè¿°ããŸããå
ãšãªãããŒã¿ã¯ data/ ã®ãã©ã«ãã«é
眮ããŠãããããèªèº«ã®ç°å¢ã«åãããŠé©åã«èšå®ããŠãã ããã åããŒãã«ã€ããŠã®IDå²ãåœãŠãšåããŒãIDã«ãã COOåœ¢åŒ 10 ã§ã®éåãªã¹ãã«ãã£ãŠãšããžã衚çŸããŠã°ã©ãããŒã¿ãå®çŸ©ããŸãã import os import numpy as np import pandas as pd from tqdm import tqdm from sklearn.preprocessing import StandardScaler from sklearn.metrics import roc_curve, roc_auc_score import matplotlib.pyplot as plt import torch import torch.nn.functional as F from torch import Tensor from torch.nn import Module import torch_geometric import torch_geometric.transforms as T from torch_geometric.nn import SAGEConv, to_hetero from torch_geometric.data import HeteroData from torch_geometric.loader import LinkNeighborLoader # ããŒã¿ã®èªã¿èŸŒã¿ïŒpandasïŒ df_recipes = pd.read_csv( '../data/food/recipes.csv' ) df_reviews = pd.read_csv( '../data/food/reviews.csv' ) # ããŒã¿æºå df_reviews = df_reviews[df_reviews.RecipeId.isin(df_recipes[ "RecipeId" ].unique())] # äžèŠããŒã¿é€å€ df_recipes[ 'RecipeServings' ] = df_recipes[ 'RecipeServings' ].fillna(df_recipes[ 'RecipeServings' ].median()) # æ¬ æå€è£å® # ãŠãŒã¶ããŒããšã¬ã·ãããŒãã®IDãããäœæ unique_user_id = df_reviews[ "AuthorId" ].unique() unique_user_id = pd.DataFrame( data={ "user_id" : unique_user_id, "mappedID" : pd.RangeIndex( len (unique_user_id)), } ) unique_recipe_id = df_reviews[ "RecipeId" ].unique() unique_recipe_id = pd.DataFrame( data={ "recipe_id" : unique_recipe_id, "mappedID" : pd.RangeIndex( len (unique_recipe_id)), } ) review_user_id = pd.merge( df_reviews[ "AuthorId" ], unique_user_id, left_on= "AuthorId" , right_on= "user_id" , how= "left" , ) review_recipe_id = pd.merge( df_reviews[ "RecipeId" ], unique_recipe_id, left_on= "RecipeId" , right_on= "recipe_id" , how= "left" , ) # ãŠãŒã¶IDãšã¬ã·ãIDã®ãšããžæ
å ±ãTensorãžå€æ tensor_review_user_id = torch.from_numpy(review_user_id[ "mappedID" ].values) tensor_review_recipe_id = torch.from_numpy(review_recipe_id[ "mappedID" ].values) tensor_edge_index_user_to_recipe = torch.stack( [tensor_review_user_id, tensor_review_recipe_id], dim= 0 , ) ååŠç ã¬ã·ãããŒããæã€ç¹åŸŽéã®ååŠçãšããŠæšæºåãè¡ããŸãã ä»åå©çšããç¹åŸŽéã¯æ¢åã®ããŒã¿ã»ããäžã«å«ãŸãã該åœã¬ã·ãã®ç³åãæ²¹åãšãã£ãæçã«ãããæ§æèŠçŽ ã®ã¿ããã©ã¡ãŒã¿ãšããŠå©çšããŸããæ¬æ¥ã§ããã°ãã®ãã§ãŒãºã§ç¹åŸŽéãšã³ãžãã¢ãªã³ã°ãªã©ãè¡ããŸããä»åã¯æ¬é¡ããå€ããŠããŸãã®ã§ã¹ãããããŸãã # ã¬ã·ãããŒãã®ç¹åŸŽéå®çŸ© recipe_feature_cols = [ "Calories" , "FatContent" , "SaturatedFatContent" , "CholesterolContent" , "SodiumContent" , "CarbohydrateContent" , "FiberContent" , "SugarContent" , "ProteinContent" , "RecipeServings" , ] df_recipes_feature = pd.merge(df_recipes, unique_recipe_id, left_on= 'RecipeId' , right_on= 'recipe_id' , how= 'left' ) df_recipes_feature = df_recipes_feature.sort_values( 'mappedID' ).set_index( 'mappedID' ) df_recipes_feature = df_recipes_feature[df_recipes_feature.index.notnull()] df_recipes_feature = df_recipes_feature[recipe_feature_cols] # æšæºå scaler = StandardScaler() scaler.fit(df_recipes_feature) scaler.transform(df_recipes_feature) df_recipes_feature = pd.DataFrame(scaler.transform(df_recipes_feature), columns=df_recipes_feature.columns) # ã¬ã·ãããŒãã®ç¹åŸŽéãTensorãžå€æ tensor_recipes_feature = torch.from_numpy(df_recipes_feature.values).to(torch.float) ããŒã¿ããŒã㌠ãããŸã§å®çŸ©ããŠããããŒã¿ãçšã㊠Pytorch ã§æ±ããããŒã¿ã»ãããšããŠããŒã¿ããŒããŒãäœæããŸãã ããŒã¿ã¯ RandomLinkSplit ãçšããŠãšããžã«å¯ŸããŠã®ããŒã¿åå²ãè¡ãããã®åŸã« LinkNeighborLoader ã§ãšããžããŒã¹ã®ããããããäœæããããŒã¿ããŒããŒãå®çŸ©ããŸãããã® LinkNeighborLoader ã§ã¯å
šãŠã®ãšããžã®äžããã©ã³ãã ãµã³ããªã³ã°ãé©çšãããã®ãšããžã®é£æ¥ããŒãããæŽã«ãµã³ããªã³ã°ãè¡ãããšã§ãå
šãŠã®ããŒãã䜿ã£ããµãã°ã©ãã«ãããããããããŒã¿äœæã宿œããŠããŸãã # HeteroDataãªããžã§ã¯ãã®äœæ data = HeteroData() data[ 'user' ].node_id = torch.arange( len (unique_user_id)) data[ 'recipe' ].node_id = torch.arange( len (unique_recipe_id)) data[ 'recipe' ].x = tensor_recipes_feature data[ 'user' , 'review' , 'recipe' ].edge_index = tensor_edge_index_user_to_recipe data = T.ToUndirected()(data) # åŠç¿ã»è©äŸ¡çšã®ããŒã¿åå² transform = T.RandomLinkSplit( num_val= 0.1 , num_test= 0.1 , disjoint_train_ratio= 0.3 , neg_sampling_ratio= 2 , add_negative_train_samples= False , edge_types=( "user" , "review" , "recipe" ), rev_edge_types=( "recipe" , "rev_review" , "user" ), ) train_data, val_data, test_data=transform(data) # åŠç¿çšããŒã¿ããŒããŒå®çŸ© edge_label_index = train_data[ "user" , "review" , "recipe" ].edge_label_index edge_label = train_data[ "user" , "review" , "recipe" ].edge_label train_loader = LinkNeighborLoader( data=train_data, num_neighbors=[ 20 , 10 ], neg_sampling_ratio= 2 , edge_label_index=(( "user" , "review" , "recipe" ), edge_label_index), edge_label=edge_label, batch_size= 256 , shuffle= True , ) # æ€èšŒçšããŒã¿ããŒããŒå®çŸ© edge_label_index = val_data[ "user" , "review" , "recipe" ].edge_label_index edge_label = val_data[ "user" , "review" , "recipe" ].edge_label val_loader = LinkNeighborLoader( data=val_data, num_neighbors=[ 20 , 10 ], edge_label_index=(( "user" , "review" , "recipe" ), edge_label_index), edge_label=edge_label, batch_size= 3 * 256 , shuffle= False , ) ã¢ãã«åŠç¿ ã¢ãã«å®çŸ© ã¢ãã«å
šäœåã®ç°¡åãªã¢ãŒããã¯ãã£ã説æãããšãåãã«ãŠãŒã¶ããŒããšã¬ã·ãããŒããåæ£è¡šçŸã«å€æããåŸã§ã2局㮠GNN ã¬ã€ã€ãŒã«ãŠåæ£è¡šçŸããéèŠç¹åŸŽéãæœåºããŠãããæåŸã«ç°ãªãããŒãéã®ãšããžååšç¢ºçãåºåãããããªã¢ãã«ãšãªã£ãŠããŸãããŸããGNN ã¬ã€ã€ãŒã«ãããã¢ã«ãŽãªãºã ã«ã¯ GraphSAGE ãçšããŠãããããã«ãã inductive ãªåé¡èšå®ã«å¯Ÿå¿ããã¢ãã«ãšãªãããã«é
æ
®ããŠããŸãã class GNN (Module): def __init__ (self, hidden_channels: int ): super ().__init__() self.conv1 = SAGEConv(hidden_channels, hidden_channels) self.conv2 = SAGEConv(hidden_channels, hidden_channels) def forward (self, x: Tensor, edge_index: Tensor) -> Tensor: x = self.conv1(x, edge_index).relu() x = self.conv2(x, edge_index) return x class Classifier (Module): def forward ( self, x_user: Tensor, x_recipe: Tensor, edge_label_index: Tensor ) -> Tensor: edge_feat_user = x_user[edge_label_index[ 0 ]] edge_feat_recipe = x_recipe[edge_label_index[ 1 ]] return (edge_feat_user * edge_feat_recipe).sum(dim=- 1 ) class Model (Module): def __init__ (self, hidden_channels: int ): super ().__init__() self.recipe_lin = torch.nn.Linear( 10 , hidden_channels) self.user_emb = torch.nn.Embedding(data[ "user" ].num_nodes, hidden_channels) self.recipe_emb = torch.nn.Embedding(data[ "recipe" ].num_nodes, hidden_channels) self.gnn = GNN(hidden_channels) self.gnn = to_hetero(self.gnn, metadata=data.metadata()) self.classifier = Classifier() def forward (self, data: HeteroData) -> Tensor: x_dict = { "user" : self.user_emb(data[ "user" ].node_id), "recipe" : self.recipe_lin(data[ "recipe" ].x) + self.recipe_emb(data[ "recipe" ].node_id), } x_dict = self.gnn(x_dict, data.edge_index_dict) pred = self.classifier( x_dict[ "user" ], x_dict[ "recipe" ], data[ "user" , "review" , "recipe" ].edge_label_index, ) return pred åŠç¿ãšè©äŸ¡ åŠç¿çšãšæ€èšŒçšã®ããŒã¿ããŒããŒãçšããŠã¢ãã«åŠç¿ãšãã®ã¢ãã«ã®è©äŸ¡ãè¡ãªã£ãŠãããŸãã è©äŸ¡ã¯ãªã³ã¯äºæž¬ã«ãããšããžãããããªããã®2å€åé¡ãšãªããã ROC-AUC 11 ã§ç²ŸåºŠã確èªããŠã¿ãããšæããŸãã def train (model, loader, device, optimizer, epoch): model.train() for epoch in range ( 1 , epoch): total_loss = total_samples = 0 for batch_data in tqdm(loader): optimizer.zero_grad() batch_data = batch_data.to(device) pred = model(batch_data) loss = F.binary_cross_entropy_with_logits( pred, batch_data[ "user" , "review" , "recipe" ].edge_label ) loss.backward() optimizer.step() total_loss += float (loss) * pred.numel() total_samples += pred.numel() print (f "Epoch: {epoch:04d}, Loss: {total_loss / total_samples:.4f}" ) def validation (model, loader, device, optimizer): y_preds = [] y_trues = [] model.eval() for batch_data in tqdm(loader): with torch.no_grad(): batch_data = batch_data.to(device) pred = model(batch_data) y_preds.append(pred) y_trues.append(batch_data[ "user" , "review" , "recipe" ].edge_label) y_pred = torch.cat(y_preds, dim= 0 ).cpu().numpy() y_true = torch.cat(y_trues, dim= 0 ).cpu().numpy() auc = roc_auc_score(y_true, y_pred) return auc, y_pred, y_true # ãã©ã¡ãŒã¿ã»ãã model = Model(hidden_channels= 64 ) device = torch.device( "cuda" if torch.cuda.is_available() else "cpu" ) optimizer = torch.optim.Adam(model.parameters(), lr= 0.001 ) model = model.to(device) # åŠç¿ã»è©äŸ¡ train(model, train_loader, device, optimizer, 6 ) auc, y_pred, y_true = validation(model, val_loader, device, optimizer) # 粟床確èªïŒROC-AUCæ²ç·ïŒ fpr, tpr, thresholds = roc_curve(y_true, y_pred) plt.plot(fpr, tpr, label=f "AUC: {auc:.3f}" ) plt.xlabel( 'FPR: False positive rate' ) plt.ylabel( 'TPR: True positive rate' ) plt.legend(loc= 'lower right' ) plt.grid() äœæããã¢ãã«ãçšããŠãæ€èšŒçšããŒã¿ãã ROC-AUC ã確èªããŸããã çµæãšããŠROC-AUCã 0.974 ãšããæ°å€ã§ããã®ã§ãæ¯èŒçé«ç²ŸåºŠã®äºæž¬ãè¡ããã¢ãã«ãšãªããŸããã ãŸããå®éçãªç²ŸåºŠææšã«å ããŠãäºæž¬ã¢ãã«ã䜿ã£ãŠãããŠãŒã¶ããŒãïŒç¹å®ãŠãŒã¶ïŒã«å¯ŸããŠã©ã³ãã æœåºããã¬ã·ãã«èå³ãæã€ããšãã£ããªã³ã¯äºæž¬ãã©ããããåœãŠãããŠããã®ããå¯èŠåããŠã¿ãŸããã å³ã®èŠæ¹ã¯çãäžã®ãªã¬ã³ãžè²ã®äžžã察象ã®ãŠãŒã¶ããŒãã瀺ããŠããããã®åšãã«ããç·è²ã®ããŒããèå³ãç€ºãæ£äŸã¬ã·ãããŒãã§ãã°ã¬ãŒè²ãèå³ã瀺ããªãè² äŸã¬ã·ãããŒããšãªããŸããäžæ®µã®å·Šå³ãæ£è§£ããŒã¿ã«ãããŠãŒã¶ãšã¬ã·ãã®é¢ä¿ã§ãå³å³ãäºæž¬çµæã«ãããŠãŒã¶ãšã¬ã·ãã®é¢ä¿ã§ããäžæ®µã®å³ã¯ãããã®æ£è§£ããŒã¿ã®å³ãšäºæž¬çµæã®å³ããæ£è§£ãšäºæž¬ãäžèŽããããŒããç·è²ãäžäžèŽã®ããŒããèµ€è²ã§è¡šçŸããå³ã§ãããããèŠããšèå³ã瀺ãã¹ãã¬ã·ãã«å¯ŸããŠãèå³ããªããšå€å¥ããŠããäºæž¬ãããã€ããããŸãããæŠãäºæž¬ãããŸãã§ããŠããããšãèŠãŠåããŸããã çµããã« ä»å㯠Heterogeneous Graph ã®ç޹ä»ãã GNN ã§ã®ã¢ããªã³ã°æ¹æ³ã«ã€ããŠç޹ä»ããŸãããã°ã©ãããŒã¿ã¯ã¡ãã£ãšçãããããæ±ãã¥ããéšåããããŸãããPytorch-geometric ãªã©ã®ãã¬ãŒã ã¯ãŒã¯ãçšããããšã§ããçšåºŠç°¡åã«å®è£
ã§ããããã«ãªã£ãŠããŸããGNN ãæ±ããããã«ãªããšåé¡è§£æ±ºã®å¹
ãåºãã£ãŠããããšæããŸãã®ã§ãèå³ãããæ¹ã¯æ¯é詊ããŠã¿ãŠãã ããã ã¢ããã³ãã«ã¬ã³ããŒãçµç€ã§ããæåŸãŸã§æ¥œããã§ãã£ãŠãã ããïŒ https://www.ntt.com/about-us/press-releases/news/article/2020/0423.html ↩ M.Schlichtkrull, T.N.Kipf, P.Bloem, R.V.D.Berg, I.Titov, and M.Welling, " Modeling relational data with graph convolutional networks ", in European Semantic Web Conference, 2018. ↩ W.Hamilton, Z.Ying, and J.Leskovec, " Inductive representation learning on large graphs ", in NeurIPS, 2017. ↩ https://distill.pub/ ↩ Wang, Xiao, et al. " Heterogeneous graph attention network. " The world wide web conference. 2019. ↩ Ali, Mehdi, et al. " Improving Inductive Link Prediction Using Hyper-relational Facts. " International Semantic Web Conference. Springer, Cham, 2021. ↩ SONG, J. AND YU, K., 2021. " Framework for Indoor Elements Classification via Inductive Learning on Floor Plan Graphs. " ISPRS International Journal of Geo-Information, Volume 10. ↩ https://www.food.com/ ↩ Pytorch-geometric以å€ã«ã DGL ã Pytorch-BigGraph ãªã©ã®ã©ã€ãã©ãªããããŸãã ↩ çè¡åã衚çŸããæ ŒçŽæ¹åŒã®1ã€ã§ãåã»è¡ã»ããŒã¿ã®3ã€ã®1次å
é
åã«ããçè¡åã衚çŸããããŒã¿åœ¢åŒã§ãã ↩ ROC-AUC ã¯äºå€åé¡ã®ã¿ã¹ã¯ã«å¯Ÿããè©äŸ¡ææšã®1ã€ãç¯å²ãšã㊠0.0 ã 1.0 ã®å€ããšãã1.0 ã«è¿ã¥ãã»ã©äºæž¬ç²ŸåºŠãé«ãããšã瀺ãã ↩