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

如何用TensorFlow tf.feature_column.crossed_column实现字符串列表特征交叉

Got it, let’s walk through how to cross your post_tags and user_tags features using TensorFlow’s tf.feature_column.crossed_column. Since both are space-separated padded strings (with actual tags plus that xyzpadxyz filler), we need to handle a few key steps first to make the cross useful.

Step 1: Parse and Clean the Padded String Features

First, we need to split the input strings into individual tags and filter out the useless xyzpadxyz padding. This ensures we only cross valid tag pairs.

import tensorflow as tf

# Example input data
sample_post = tf.constant(["the-oscars brad-pitt xyzpadxyz xyzpadxyz"])
sample_user = tf.constant(["brad-pitt universal sag the-academy-awards xyzpadxyz xyzpadxyz"])

# Split strings into ragged tensors (handles variable-length tags)
post_tags_split = tf.strings.split(sample_post, sep=" ").to_ragged()
user_tags_split = tf.strings.split(sample_user, sep=" ").to_ragged()

# Filter out the padding value
post_tags_clean = tf.ragged.boolean_mask(post_tags_split, post_tags_split != "xyzpadxyz")
user_tags_clean = tf.ragged.boolean_mask(user_tags_split, user_tags_split != "xyzpadxyz")
Step 2: Convert Cleaned Tags to Categorical Columns

Since tag vocabularies are usually large, hash bucket categorical columns are the most practical choice (no need to predefine every possible tag). If you have a fixed vocabulary list, you can use categorical_column_with_vocabulary_list instead for zero collision risk.

# Create hash bucket columns for each feature
post_tag_cat_col = tf.feature_column.categorical_column_with_hash_bucket(
    key="post_tags",
    hash_bucket_size=1000  # Adjust based on your expected unique tag count
)

user_tag_cat_col = tf.feature_column.categorical_column_with_hash_bucket(
    key="user_tags",
    hash_bucket_size=1000
)
Step 3: Create the Crossed Column

Now we can cross the two categorical columns. This will generate all possible pairwise combinations between the tags in post_tags and user_tags (e.g., ("the-oscars", "brad-pitt"), ("brad-pitt", "universal"), etc.).

# Cross the two tag columns
crossed_tags_col = tf.feature_column.crossed_column(
    [post_tag_cat_col, user_tag_cat_col],
    hash_bucket_size=10000  # Use a larger bucket size for cross combinations
)
Step 4: Convert Crossed Column to a Model-Ready Format

Categorical and crossed columns can’t be fed directly into a model—we need to convert them to either an embedding column (great for high-cardinality data) or an indicator column (only feasible if your hash bucket size is small).

# Convert to embedding column (ideal for most cases)
crossed_tags_embedding = tf.feature_column.embedding_column(
    crossed_tags_col,
    dimension=32  # Adjust based on your model's needs
)

# Alternative: Indicator column (only use if hash_bucket_size is <10k)
# crossed_tags_indicator = tf.feature_column.indicator_column(crossed_tags_col)
Step 5: Integrate into a Keras Model

Finally, we’ll wire this into a Keras model. We use a Lambda layer to clean the input strings on the fly, then pass the processed features to a DenseFeatures layer to handle the crossed column.

# Define input layers
post_input = tf.keras.layers.Input(shape=(1,), dtype=tf.string, name="post_tags")
user_input = tf.keras.layers.Input(shape=(1,), dtype=tf.string, name="user_tags")

# Helper function to clean tags and convert to sparse tensor (required for feature columns)
def clean_and_convert_to_sparse(input_tensor):
    split = tf.strings.split(input_tensor, sep=" ").to_ragged()
    cleaned = tf.ragged.boolean_mask(split, split != "xyzpadxyz")
    return cleaned.to_sparse()

# Process inputs
post_processed = tf.keras.layers.Lambda(clean_and_convert_to_sparse)(post_input)
user_processed = tf.keras.layers.Lambda(clean_and_convert_to_sparse)(user_input)

# Create feature layer
feature_layer = tf.keras.layers.DenseFeatures([crossed_tags_embedding])
crossed_features = feature_layer({
    "post_tags": post_processed,
    "user_tags": user_processed
})

# Build the rest of the model
dense_layer = tf.keras.layers.Dense(64, activation="relu")(crossed_features)
output = tf.keras.layers.Dense(1, activation="sigmoid")(dense_layer)

model = tf.keras.Model(inputs=[post_input, user_input], outputs=output)
model.compile(optimizer="adam", loss="binary_crossentropy", metrics=["accuracy"])
Key Notes to Remember
  • Always Filter Padding: Leaving xyzpadxyz in will create useless cross pairs that add noise to your model—don’t skip this step!
  • Hash Bucket Sizing: Adjust hash_bucket_size based on your dataset. Too small and you’ll get hash collisions; too large and you’ll waste memory.
  • Multi-Valued Behavior: crossed_column automatically handles multi-valued features by generating all possible tag pairs, which is perfect for matching post tags to a user’s browsing history tags.

内容的提问来源于stack exchange,提问作者andrewm4894

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.25 07:05:32