TensorFlow简易Estimator输入函数构建问题求助
Hey there! Let's work through this problem together. It sounds like you might be running into issues with either manually defining your labels (which can lead to typos or mismatches) or ensuring your input function correctly ties features to labels. Here's a structured solution:
1. Automate Label Generation (Avoid Manual Errors)
Instead of hardcoding your labels array (which is error-prone, especially with larger datasets), use numpy's vectorized operations to generate labels based on your rules automatically. This ensures every number gets the correct label without mistakes.
import numpy as np nmbrs = np.array([10., 1., 2., 3., 4., 5., 6. , 7., 8., 9.]) # Generate labels using conditional logic labels = np.where(nmbrs < 5, 0., np.where(nmbrs == 5, 5., 10.)) print(labels) # Output: [10. 0. 0. 0. 0. 5. 10. 10. 10. 10.]
This code uses nested np.where() calls to apply your exact label rules:
- If
x < 5: assign0. - Else if
x == 5: assign5. - Else (x >5): assign
10.
2. Validate Your Input Function
Once your labels are correctly generated, double-check your input function setup. Make sure you've defined batch_size (a common oversight) and that your data types are consistent (all floats here, which is good).
Here's the complete, corrected input function code:
import tensorflow as tf batch_size = 2 # Adjust this value to match your training needs input_fn = tf.estimator.inputs.numpy_input_fn( x={'numbers': nmbrs}, # No need to re-wrap nmbrs since it's already a numpy array y=labels, batch_size=batch_size, num_epochs=None, shuffle=True ) # Test the input function to verify batches for batch in input_fn(): print("Feature batch:", batch[0]['numbers'].numpy()) print("Label batch:", batch[1].numpy()) break # Stop after first batch to check output
3. Common Pitfalls to Watch For
- Undefined
batch_size: Always make surebatch_sizeis defined before using it in the input function. - Mismatched feature/label lengths: Ensure
nmbrsandlabelshave the same number of elements (the automated label generation takes care of this). - Data type inconsistencies: Keep your features and labels in the same dtype (e.g., all floats) to avoid TensorFlow errors.
内容的提问来源于stack exchange,提问作者cross wade

