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
« 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
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
10"""hopwise.data.interaction
11############################
12"""
14from collections.abc import Mapping
16import numpy as np
17import pandas as pd
18import torch
19import torch.nn.utils.rnn as rnn_utils
22def _convert_to_tensor(data):
23 """This function can convert common data types (list, pandas.Series, numpy.ndarray, torch.Tensor) into torch.Tensor.
25 Args:
26 data (list, pandas.Series, numpy.ndarray, torch.Tensor): Origin data.
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
44class Interaction(Mapping):
45 """The basic class representing a batch of interaction records.
47 Note:
48 While training, there is no strict rules for data in one Interaction object.
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.
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 ======= ======= =======
67 Some wrong examples for Interaction objects used in testing:
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
76 occur earlier than this user's negative cases
77 1 3 1
78 1 1 0
79 2 3 1
80 ... ... ...
81 ======= ======= ======= ============
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 ======= ======= ======= ========
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 """
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])
120 def __iter__(self):
121 return self.interaction.__iter__()
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}'")
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()
136 ret = {}
137 for k in self.interaction:
138 ret[k] = self.interaction[k][index]
139 return Interaction(ret)
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
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]
151 def __contains__(self, item):
152 return item in self.interaction
154 def __len__(self):
155 return self.length
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)
166 def __repr__(self):
167 return self.__str__()
169 @property
170 def columns(self):
171 """Returns:
172 list of str: The columns of interaction.
173 """
174 return list(self.interaction.keys())
176 def size(self, dim=0):
177 """Get the size of the interaction along a specific dimension.
179 Args:
180 dim (int): The dimension to get the size of. Default is 0 (batch dimension).
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
190 def to(self, device, selected_field=None):
191 """Transfer Tensors in this Interaction object to the specified device.
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.
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]
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)
218 def cpu(self):
219 """Transfer Tensors in this Interaction object to cpu.
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)
229 def numpy(self):
230 """Transfer Tensors to numpy arrays.
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
241 def repeat(self, sizes):
242 """Repeats each tensor along the batch dim.
244 Args:
245 sizes (int): repeat times.
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
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
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)
266 def repeat_interleave(self, repeats, dim=0):
267 """Similar to repeat_interleave of PyTorch.
269 Details can be found in:
271 https://pytorch.org/docs/stable/tensors.html?highlight=repeat#torch.Tensor.repeat_interleave
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)
281 def update(self, new_inter):
282 """Similar to ``dict.update()``
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]
290 def drop(self, column):
291 """Drop column in interaction.
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]
300 def _reindex(self, index):
301 """Reset the index of interaction inplace.
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]
309 def shuffle(self):
310 """Shuffle current interaction inplace."""
311 index = torch.randperm(self.length)
312 self._reindex(index)
314 def sort(self, by, ascending=True):
315 """Sort the current interaction inplace.
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}].")
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}].")
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.")
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)
358 def add_prefix(self, prefix):
359 """Add prefix to current interaction's columns.
361 Args:
362 prefix (str): The prefix to be added.
363 """
364 self.interaction = {prefix + key: value for key, value in self.interaction.items()}
367def cat_interactions(interactions):
368 """Concatenate list of interactions to single interaction.
370 Args:
371 interactions (list of :class:`Interaction`): List of interactions to be concatenated.
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.")
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.")
386 new_inter = {col: torch.cat([inter[col] for inter in interactions]) for col in columns_set}
387 return Interaction(new_inter)