追加含新标签的MultiIndex DataFrame时保留旧索引整数位置
我在做推荐系统的矩阵分解模型(LightFM)训练时,遇到了ID映射的问题,先分享下基础场景的实现,再说说扩展场景的痛点,最后给出解决方案:
基础场景
在推荐服务中,我基于用户-物品交互数据训练矩阵分解模型(LightFM)。为使模型达到最佳效果,需将用户ID和物品ID映射为从0开始的连续整数ID。
我使用pandas DataFrame处理数据,发现MultiIndex在创建该映射时非常便捷,示例如下:
ratings = [{'user_id': 1, 'item_id': 1, 'rating': 1.0}, {'user_id': 1, 'item_id': 3, 'rating': 1.0}, {'user_id': 3, 'item_id': 1, 'rating': 1.0}, {'user_id': 3, 'item_id': 3, 'rating': 1.0}] df = pd.DataFrame(ratings, columns=['user_id', 'item_id', 'rating']) df = df.set_index(['user_id', 'item_id']) df # 输出: # rating # user_id item_id # 1 1 1.0 # 3 1.0 # 3 1 1.0 # 3 1.0
之后可通过以下方式获取连续映射:
df.index.labels[0] # 对应用户 # 输出:FrozenNDArray([0, 0, 1, 1], dtype='int8') df.index.labels[1] # 对应物品 # 输出:FrozenNDArray([0, 1, 0, 1], dtype='int8')
后续可使用df.index.levels[0].get_loc方法完成反向映射,效果很好!
扩展场景
现在我尝试优化模型训练流程,希望基于新数据进行增量训练,同时保留原有的ID映射。示例如下:
new_ratings = [{'user_id': 2, 'item_id': 1, 'rating': 1.0}, {'user_id': 2, 'item_id': 2, 'rating': 1.0}] df2 = pd.DataFrame(new_ratings, columns=['user_id', 'item_id', 'rating']) df2 = df2.set_index(['user_id', 'item_id']) df2 # 输出: # rating # user_id item_id # 2 1 1.0 # 2 1.0
将新评分数据追加到旧DataFrame:
df3 = df.append(df2) df3 # 输出: # rating # user_id item_id # 1 1 1.0 # 3 1.0 # 3 1 1.0 # 3 1.0 # 2 1 1.0 # 2 1.0
看起来正常,但:
df3.index.labels[0] # 对应用户 # 输出:FrozenNDArray([0, 0, 2, 2, 1, 1], dtype='int8') df3.index.labels[1] # 对应物品 # 输出:FrozenNDArray([0, 2, 0, 2, 0, 1], dtype='int8')
我特意在新数据中加入user_id=2和item_id=2,以此说明问题:在df3中,原user_id=3和item_id=3的标签从整数位置1变为2,映射关系不再一致。我期望的用户和物品映射分别为[0, 0, 1, 1, 2, 2]和[0, 1, 0, 1, 0, 2]。
这可能是由于pandas Index对象的排序机制导致的,我不确定使用MultiIndex策略能否实现需求,寻求有效解决方法:)
补充说明:
- 我认为使用DataFrame较为便捷,但仅用MultiIndex做ID映射,不使用MultiIndex的方案也可接受。
- 无法保证新数据中的user_id和item_id大于旧数据中的现有值,因此示例中在已有[1,3]的情况下加入了ID 2。
- 增量训练时需存储ID映射,若仅加载部分新数据,需存储旧DataFrame和ID映射,最好能统一存储(如索引或列)。
- 编辑补充:新增需求:允许原DataFrame行重排(如存在重复评分时保留最新记录)。
解决方案(感谢@jpp的原始方案)
我对@jpp的方案进行了修改,以满足后续新增的需求(标记为EDIT)。该方案完全符合标题中的原始需求,无论原DataFrame行如何重排,都能保留旧索引的整数位置。我还将代码封装为函数:
from itertools import chain from toolz import unique def expand_index(source, target, index_cols=['user_id', 'item_id']): # Elevate index to series, keeping source with index temp = source.reset_index() target = target.reset_index() # Convert columns to categorical, using the source index and target columns for col in index_cols: i = source.index.names.index(col) col_cats = list(unique(chain(source.index.levels[i], target[col]))) temp[col] = pd.Categorical(temp[col], categories=col_cats) target[col] = pd.Categorical(target[col], categories=col_cats) # Convert series back to index source = temp.set_index(index_cols) target = target.set_index(index_cols) return source, target def concat_expand_index(old, new): old, new = expand_index(old, new) return pd.concat([old, new]) df3 = concat_expand_index(df, df2)
结果:
df3.index.labels[0] # 对应用户 # 输出:FrozenNDArray([0, 0, 1, 1, 2, 2], dtype='int8') df3.index.labels[1] # 对应物品 # 输出:FrozenNDArray([0, 1, 0, 1, 0, 2], dtype='int8')
内容的提问来源于stack exchange,提问作者Fulco
相关产品推荐
相关产品推荐

