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

使用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

  1. Incomplete generator function: Your input_generator_fn is truncated (for res in res...), which would throw a syntax error before even reaching the TensorFlow dataset logic.
  2. Missing output signature: from_generator requires 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.
  3. Potential tensor creation conflicts: Using tf.constant directly 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

  1. Completed generator logic: Added a full loop over resolutions and included while True to create an infinite data stream—perfect for continuous inference performance testing, as you won't run out of data mid-test.
  2. Numpy-based image generation: Switched from tf.constant to numpy array creation + tf.convert_to_tensor. This avoids potential graph/eager mode conflicts that can arise when creating tensors directly inside a generator.
  3. Explicit output signature: Used tf.TensorSpec to define the exact shape and dtype of the source/target tensors. This is the most common fix for InvalidArgumentError with from_generator—TensorFlow needs clear instructions to build a valid dataset pipeline.
  4. 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 and size[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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.22 08:25:33