如何用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.
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")
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 )
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 )
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)
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"])
- Always Filter Padding: Leaving
xyzpadxyzin will create useless cross pairs that add noise to your model—don’t skip this step! - Hash Bucket Sizing: Adjust
hash_bucket_sizebased on your dataset. Too small and you’ll get hash collisions; too large and you’ll waste memory. - Multi-Valued Behavior:
crossed_columnautomatically 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

