You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

PyTorch实现Keras风格哈希交叉积转换的技术疑问

在PyTorch中实现Keras的哈希交叉积转换

Keras中的HashedCrossing示例

>>> layer = keras.layers.HashedCrossing(num_bins=5, output_mode='one_hot')
>>> feat1 = np.array([1, 5, 2, 1, 4])
>>> feat2 = np.array([2, 9, 42, 37, 8])
>>> layer((feat1, feat2))
<tf.Tensor: shape=(5, 5), dtype=float32, numpy=
array([[0., 0., 1., 0., 0.],
       [1., 0., 0., 0., 0.],
       [0., 0., 0., 0., 1.],
       [1., 0., 0., 0., 0.],
       [0., 0., 1., 0., 0.]], dtype=float32)>
>>> layer2 = keras.layers.HashedCrossing(num_bins=5, output_mode='int')
>>> layer2((feat1, feat2))
<tf.Tensor: shape=(5,), dtype=int64, numpy=array([2, 0, 4, 0, 2])>

官方说明:

该层使用“哈希技巧”对类别特征进行交叉转换,概念上可理解为:hash(concatenate(features)) % num_bins。

问题解析与你的尝试

你对concatenate(features)的疑问:确实需要对每个特征“对”进行哈希——这里的concatenate指的是把每个样本对应的一组特征(比如feat1[i]和feat2[i])拼接成一个整体,再对这个整体计算哈希值,而非对整个特征张量做拼接哈希。

你尝试的实现:

>>> cross_product_idx = (feat1*feat2.max()+1 + feat2) % num_bins
>>> cross_product = nn.functional.one_hot(cross_product_idx, num_bins)

这个实现能运行,但确实存在分布问题:这种方式是手动构造特征对的唯一标识,若feat1、feat2的数值分布不均,计算出的索引会集中在某些区间,导致哈希桶负载不均,而哈希函数的核心作用就是把离散的特征对尽可能均匀地映射到各个桶中。

PyTorch实现方案

基于Torch原生哈希的高效实现

import torch
import torch.nn as nn

def hashed_crossing(feat1, feat2, num_bins, output_mode="one_hot"):
    # 将每个样本的特征对拼接为二维张量(shape: [N, 2])
    combined = torch.stack([feat1, feat2], dim=-1)
    # 计算每个特征对的哈希值(torch.hash支持张量元素级哈希)
    hashed_vals = torch.hash(combined)
    # 处理负哈希值,取模得到桶索引
    bin_indices = torch.abs(hashed_vals) % num_bins
    
    if output_mode == "int":
        return bin_indices
    elif output_mode == "one_hot":
        # 转换为float类型对齐Keras输出
        return nn.functional.one_hot(bin_indices, num_classes=num_bins).float()
    else:
        raise ValueError(f"不支持的输出模式: {output_mode}")

# 测试用例
feat1 = torch.tensor([1, 5, 2, 1, 4])
feat2 = torch.tensor([2, 9, 42, 37, 8])
num_bins = 5

# 验证one_hot模式
print("one_hot输出:")
print(hashed_crossing(feat1, feat2, num_bins, output_mode="one_hot"))

# 验证int模式
print("\nint输出:")
print(hashed_crossing(feat1, feat2, num_bins, output_mode="int"))

兼容旧版本Torch的实现(基于Python哈希)

如果你的Torch版本不支持torch.hash(),可以用Python内置哈希函数处理特征对:

def hashed_crossing_compat(feat1, feat2, num_bins, output_mode="one_hot"):
    bin_indices = []
    # 遍历每个特征对计算哈希
    for f1, f2 in zip(feat1.numpy(), feat2.numpy()):
        hash_val = hash((f1, f2))
        bin_idx = abs(hash_val) % num_bins
        bin_indices.append(bin_idx)
    bin_indices = torch.tensor(bin_indices)
    
    if output_mode == "int":
        return bin_indices
    elif output_mode == "one_hot":
        return nn.functional.one_hot(bin_indices, num_classes=num_bins).float()

关键说明

  1. 哈希对象:必须针对每个样本的特征对独立哈希,这样才能保证每个交叉特征对应唯一的哈希值。
  2. 哈希均匀性:哈希函数能将离散的特征对均匀映射到各个桶中,避免手动构造ID带来的分布倾斜问题。

内容的提问来源于stack exchange,提问作者David Davó

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.18 00:13:18