使用TensorFlow Dataset.from_generator时遇InvalidArgumentError问题求助
Fixing InvalidArgumentError with tf.data.Dataset.from_generator for Custom Image Datasets
Let's break down what's going wrong here and fix it step by step. I see a couple of key issues in your code snippet that are almost certainly causing that InvalidArgumentError when using tf.data.Dataset.from_generator.
Key Problems in Your Original Code
- Incomplete generator function: Your
input_generator_fnis truncated (for res in res...), which would throw a syntax error before even reaching the TensorFlow dataset logic. - Missing output signature:
from_generatorrequires explicit information about the shape and dtype of the tensors your generator yields. Without this, TensorFlow can't properly construct the dataset pipeline, leading to argument errors. - Potential tensor creation conflicts: Using
tf.constantdirectly inside a generator can cause issues with TensorFlow's eager/graph mode separation, especially if the generator isn't properly aligned with eager execution.
Corrected Full Code
Here's a revised version of your code that resolves these issues and works reliably for generating a custom image dataset:
import tensorflow as tf import random import numpy as np # Define target image resolution (width, height) resolutions = [(2048, 1080)] def generate_image(size, channels): # Use numpy to generate the base image first (avoids graph mode conflicts) image_value = random.random() # Shape follows (height, width, channels) since size is (width, height) image_np = np.full(shape=(size[1], size[0], channels), fill_value=image_value, dtype=np.float32) return tf.convert_to_tensor(image_np, dtype=tf.float32) def generate_single_input(size): source = generate_image(size, 3) target = generate_image(size, 3) # Add batch dimension (matches your original [1, H, W, C] shape) source = tf.expand_dims(source, axis=0) target = tf.expand_dims(target, axis=0) return source, target def input_generator_fn(): # Loop through resolutions (add while True for infinite stream, ideal for performance testing) for res in resolutions: while True: yield generate_single_input(res) # Critical: Define the output signature so TensorFlow knows what to expect output_signature = ( tf.TensorSpec(shape=(1, 1080, 2048, 3), dtype=tf.float32), tf.TensorSpec(shape=(1, 1080, 2048, 3), dtype=tf.float32) ) # Create the dataset dataset = tf.data.Dataset.from_generator( generator=input_generator_fn, output_signature=output_signature ) # Test the dataset (verify it works without errors) for src, tgt in dataset.take(1): print(f"Source shape: {src.shape}, Target shape: {tgt.shape}")
What Changed & Why
- Completed generator logic: Added a full loop over
resolutionsand includedwhile Trueto create an infinite data stream—perfect for continuous inference performance testing, as you won't run out of data mid-test. - Numpy-based image generation: Switched from
tf.constantto numpy array creation +tf.convert_to_tensor. This avoids potential graph/eager mode conflicts that can arise when creating tensors directly inside a generator. - Explicit output signature: Used
tf.TensorSpecto define the exact shape and dtype of the source/target tensors. This is the most common fix forInvalidArgumentErrorwithfrom_generator—TensorFlow needs clear instructions to build a valid dataset pipeline. - Clearer batch dimension handling: Separated the batch dimension addition from the base image generation, making it easier to adjust batch sizes later if needed.
Additional Checks If You Still See Errors
- Double-check resolution order: Your code uses
size[1]for height andsize[0]for width (matching (2048,1080) as width x height), which is correct—just ensure this aligns with your model's input expectations. - Verify eager mode: TensorFlow 2.x defaults to eager execution, but if you're running this inside a
tf.function(graph mode), ensure your generator logic is compatible (the numpy-based approach should work here).
内容的提问来源于stack exchange,提问作者Alexey
相关产品推荐
相关产品推荐

