Coverage for hopwise/data/interaction.py: 69%

169 statements  

« prev     ^ index     » next       coverage.py v7.16.2, created at 2026-09-30 13:25 +0000

1# @Time : 2020/7/10 

2# @Author : Yupeng Hou 

3# @Email : houyupeng@ruc.edu.cn 

4 

5# UPDATE 

6# @Time : 2022/7/8, 2020/9/15, 2020/9/16, 2020/8/12 

7# @Author : Zhen Tian, Yupeng Hou, Yushuo Chen, Xingyu Pan 

8# @email : chenyuwuxinn@gmail.com, houyupeng@ruc.edu.cn, chenyushuo@ruc.edu.cn, panxy@ruc.edu.cn 

9 

10"""hopwise.data.interaction 

11############################ 

12""" 

13 

14from collections.abc import Mapping 

15 

16import numpy as np 

17import pandas as pd 

18import torch 

19import torch.nn.utils.rnn as rnn_utils 

20 

21 

22def _convert_to_tensor(data): 

23 """This function can convert common data types (list, pandas.Series, numpy.ndarray, torch.Tensor) into torch.Tensor. 

24 

25 Args: 

26 data (list, pandas.Series, numpy.ndarray, torch.Tensor): Origin data. 

27 

28 Returns: 

29 torch.Tensor: Converted tensor from `data`. 

30 """ # noqa: E501 

31 elem = data[0] 

32 if isinstance(elem, (float, int, np.float64, np.int64)): 

33 new_data = torch.as_tensor(data) 

34 elif isinstance(elem, (list, tuple, pd.Series, np.ndarray, torch.Tensor)): 

35 seq_data = [torch.as_tensor(d) for d in data] 

36 new_data = rnn_utils.pad_sequence(seq_data, batch_first=True) 

37 else: 

38 raise ValueError(f"[{type(elem)}] is not supported!") 

39 if new_data.dtype == torch.float64: 

40 new_data = new_data.float() 

41 return new_data 

42 

43 

44class Interaction(Mapping): 

45 """The basic class representing a batch of interaction records. 

46 

47 Note: 

48 While training, there is no strict rules for data in one Interaction object. 

49 

50 While testing, it should be guaranteed that all interaction records of one single 

51 user will not appear in different Interaction object, and records of the same user 

52 should be continuous. Meanwhile, the positive cases of one user always need to occur 

53 **earlier** than this user's negative cases. 

54 

55 A correct example: 

56 ======= ======= ======= 

57 user_id item_id label 

58 ======= ======= ======= 

59 1 2 1 

60 1 6 1 

61 1 3 1 

62 1 1 0 

63 2 3 1 

64 ... ... ... 

65 ======= ======= ======= 

66 

67 Some wrong examples for Interaction objects used in testing: 

68 

69 1. 

70 ======= ======= ======= ============ 

71 user_id item_id label 

72 ======= ======= ======= ============ 

73 1 2 1 

74 1 6 0 # positive cases of one user always need to 

75 

76 occur earlier than this user's negative cases 

77 1 3 1 

78 1 1 0 

79 2 3 1 

80 ... ... ... 

81 ======= ======= ======= ============ 

82 

83 2. 

84 ======= ======= ======= ======== 

85 user_id item_id label 

86 ======= ======= ======= ======== 

87 1 2 1 

88 1 6 1 

89 1 3 1 

90 2 3 1 # records of the same user should be continuous. 

91 1 1 0 

92 ... ... ... 

93 ======= ======= ======= ======== 

94 

95 Attributes: 

96 interaction (dict or pandas.DataFrame): keys are meaningful str (also can be called field name), 

97 and values are Torch Tensor of numpy Array with shape (batch_size, \\*). 

98 """ 

99 

100 def __init__(self, interaction): 

101 self.interaction = dict() 

102 if isinstance(interaction, dict): 

103 for key, value in interaction.items(): 

104 if isinstance(value, (list, np.ndarray)): 

105 self.interaction[key] = _convert_to_tensor(value) 

106 elif isinstance(value, torch.Tensor): 

107 self.interaction[key] = value 

108 else: 

109 raise ValueError(f"The type of {key}[{type(value)}] is not supported!") 

110 elif isinstance(interaction, pd.DataFrame): 

111 for key in interaction: 

112 value = interaction[key].values 

113 self.interaction[key] = _convert_to_tensor(value) 

114 else: 

115 raise ValueError(f"[{type(interaction)}] is not supported for initialize `Interaction`!") 

116 self.length = -1 

117 for k in self.interaction: 

118 self.length = max(self.length, self.interaction[k].unsqueeze(-1).shape[0]) 

119 

120 def __iter__(self): 

121 return self.interaction.__iter__() 

122 

123 def __getattr__(self, item): 

124 if "interaction" not in self.__dict__: 

125 raise AttributeError("'Interaction' object has no attribute 'interaction'") 

126 if item in self.interaction: 

127 return self.interaction[item] 

128 raise AttributeError(f"'Interaction' object has no attribute '{item}'") 

129 

130 def __getitem__(self, index): 

131 if isinstance(index, str): 

132 return self.interaction[index] 

133 if isinstance(index, (np.ndarray, torch.Tensor)): 

134 index = index.tolist() 

135 

136 ret = {} 

137 for k in self.interaction: 

138 ret[k] = self.interaction[k][index] 

139 return Interaction(ret) 

140 

141 def __setitem__(self, key, value): 

142 if not isinstance(key, str): 

143 raise KeyError(f"{type(key)} object does not support item assigment") 

144 self.interaction[key] = value 

145 

146 def __delitem__(self, key): 

147 if key not in self.interaction: 

148 raise KeyError(f"{type(key)} object does not in this interaction") 

149 del self.interaction[key] 

150 

151 def __contains__(self, item): 

152 return item in self.interaction 

153 

154 def __len__(self): 

155 return self.length 

156 

157 def __str__(self): 

158 info = [f"The batch_size of interaction: {self.length}"] 

159 for k in self.interaction: 

160 inter = self.interaction[k] 

161 temp_str = f" {k}, {inter.shape}, {inter.device.type}, {inter.dtype}" 

162 info.append(temp_str) 

163 info.append("\n") 

164 return "\n".join(info) 

165 

166 def __repr__(self): 

167 return self.__str__() 

168 

169 @property 

170 def columns(self): 

171 """Returns: 

172 list of str: The columns of interaction. 

173 """ 

174 return list(self.interaction.keys()) 

175 

176 def size(self, dim=0): 

177 """Get the size of the interaction along a specific dimension. 

178 

179 Args: 

180 dim (int): The dimension to get the size of. Default is 0 (batch dimension). 

181 

182 Returns: 

183 int: The size of the interaction along the specified dimension. 

184 """ 

185 size = 0 

186 for k in self.interaction: 

187 size = max(size, self.interaction[k].size(dim)) 

188 return size 

189 

190 def to(self, device, selected_field=None): 

191 """Transfer Tensors in this Interaction object to the specified device. 

192 

193 Args: 

194 device (torch.device): target device. 

195 selected_field (str or iterable object, optional): if specified, only Tensors 

196 with keys in selected_field will be sent to device. 

197 

198 Returns: 

199 Interaction: a coped Interaction object with Tensors which are sent to 

200 the specified device. 

201 """ 

202 ret = {} 

203 if isinstance(selected_field, str): 

204 selected_field = [selected_field] 

205 

206 if selected_field is not None: 

207 selected_field = set(selected_field) 

208 for k in self.interaction: 

209 if k in selected_field: 

210 ret[k] = self.interaction[k].to(device) 

211 else: 

212 ret[k] = self.interaction[k] 

213 else: 

214 for k in self.interaction: 

215 ret[k] = self.interaction[k].to(device) 

216 return Interaction(ret) 

217 

218 def cpu(self): 

219 """Transfer Tensors in this Interaction object to cpu. 

220 

221 Returns: 

222 Interaction: a coped Interaction object with Tensors which are sent to cpu. 

223 """ 

224 ret = {} 

225 for k in self.interaction: 

226 ret[k] = self.interaction[k].cpu() 

227 return Interaction(ret) 

228 

229 def numpy(self): 

230 """Transfer Tensors to numpy arrays. 

231 

232 Returns: 

233 dict: keys the same as Interaction object, are values are corresponding numpy 

234 arrays transformed from Tensor. 

235 """ 

236 ret = {} 

237 for k in self.interaction: 

238 ret[k] = self.interaction[k].numpy() 

239 return ret 

240 

241 def repeat(self, sizes): 

242 """Repeats each tensor along the batch dim. 

243 

244 Args: 

245 sizes (int): repeat times. 

246 

247 Example: 

248 >>> a = Interaction({'k': torch.zeros(4)}) 

249 >>> a.repeat(3) 

250 The batch_size of interaction: 12 

251 k, torch.Size([12]), cpu 

252 

253 >>> a = Interaction({'k': torch.zeros(4, 7)}) 

254 >>> a.repeat(3) 

255 The batch_size of interaction: 12 

256 k, torch.Size([12, 7]), cpu 

257 

258 Returns: 

259 a copyed Interaction object with repeated Tensors. 

260 """ 

261 ret = {} 

262 for k in self.interaction: 

263 ret[k] = self.interaction[k].repeat([sizes] + [1] * (len(self.interaction[k].shape) - 1)) 

264 return Interaction(ret) 

265 

266 def repeat_interleave(self, repeats, dim=0): 

267 """Similar to repeat_interleave of PyTorch. 

268 

269 Details can be found in: 

270 

271 https://pytorch.org/docs/stable/tensors.html?highlight=repeat#torch.Tensor.repeat_interleave 

272 

273 Note: 

274 ``torch.repeat_interleave()`` is supported in PyTorch >= 1.2.0. 

275 """ 

276 ret = {} 

277 for k in self.interaction: 

278 ret[k] = self.interaction[k].repeat_interleave(repeats, dim=dim) 

279 return Interaction(ret) 

280 

281 def update(self, new_inter): 

282 """Similar to ``dict.update()`` 

283 

284 Args: 

285 new_inter (Interaction): current interaction will be updated by new_inter. 

286 """ 

287 for k in new_inter.interaction: 

288 self.interaction[k] = new_inter.interaction[k] 

289 

290 def drop(self, column): 

291 """Drop column in interaction. 

292 

293 Args: 

294 column (str): the column to be dropped. 

295 """ 

296 if column not in self.interaction: 

297 raise ValueError(f"Column [{column}] is not in [{self}].") 

298 del self.interaction[column] 

299 

300 def _reindex(self, index): 

301 """Reset the index of interaction inplace. 

302 

303 Args: 

304 index: the new index of current interaction. 

305 """ 

306 for k in self.interaction: 

307 self.interaction[k] = self.interaction[k][index] 

308 

309 def shuffle(self): 

310 """Shuffle current interaction inplace.""" 

311 index = torch.randperm(self.length) 

312 self._reindex(index) 

313 

314 def sort(self, by, ascending=True): 

315 """Sort the current interaction inplace. 

316 

317 Args: 

318 by (str or list of str): Field that as the key in the sorting process. 

319 ascending (bool or list of bool, optional): Results are ascending if ``True``, otherwise descending. 

320 Defaults to ``True`` 

321 """ 

322 if isinstance(by, str): 

323 if by not in self.interaction: 

324 raise ValueError(f"[{by}] is not exist in interaction [{self}].") 

325 by = [by] 

326 elif isinstance(by, (list, tuple)): 

327 for b in by: 

328 if b not in self.interaction: 

329 raise ValueError(f"[{b}] is not exist in interaction [{self}].") 

330 else: 

331 raise TypeError(f"Wrong type of by [{by}].") 

332 

333 if isinstance(ascending, bool): 

334 ascending = [ascending] 

335 elif isinstance(ascending, (list, tuple)): 

336 for a in ascending: 

337 if not isinstance(a, bool): 

338 raise TypeError(f"Wrong type of ascending [{ascending}].") 

339 else: 

340 raise TypeError(f"Wrong type of ascending [{ascending}].") 

341 

342 if len(by) != len(ascending): 

343 if len(ascending) == 1: 

344 ascending = ascending * len(by) 

345 else: 

346 raise ValueError(f"by [{by}] and ascending [{ascending}] should have same length.") 

347 

348 for b, a in zip(by[::-1], ascending[::-1]): 

349 if len(self.interaction[b].shape) == 1: 

350 key = self.interaction[b] 

351 else: 

352 key = self.interaction[b][..., 0] 

353 index = np.argsort(key, kind="stable") 

354 if not a: 

355 index = torch.tensor(np.array(index)[::-1]) 

356 self._reindex(index) 

357 

358 def add_prefix(self, prefix): 

359 """Add prefix to current interaction's columns. 

360 

361 Args: 

362 prefix (str): The prefix to be added. 

363 """ 

364 self.interaction = {prefix + key: value for key, value in self.interaction.items()} 

365 

366 

367def cat_interactions(interactions): 

368 """Concatenate list of interactions to single interaction. 

369 

370 Args: 

371 interactions (list of :class:`Interaction`): List of interactions to be concatenated. 

372 

373 Returns: 

374 :class:`Interaction`: Concatenated interaction. 

375 """ 

376 if not isinstance(interactions, (list, tuple)): 

377 raise TypeError(f"Interactions [{interactions}] should be list or tuple.") 

378 if len(interactions) == 0: 

379 raise ValueError(f"Interactions [{interactions}] should have some interactions.") 

380 

381 columns_set = set(interactions[0].columns) 

382 for inter in interactions: 

383 if columns_set != set(inter.columns): 

384 raise ValueError(f"Interactions [{interactions}] should have some interactions.") 

385 

386 new_inter = {col: torch.cat([inter[col] for inter in interactions]) for col in columns_set} 

387 return Interaction(new_inter)