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

如何在tf.data.Dataset.map中用tf.string访问字典标注值?

问题

我正在创建tf.data.Dataset,通过list_files获取所有图片路径。标注信息以JSON文件形式存储在磁盘上,结构如下:

{ 
   "img1.png": {
                "data": ...
               },
   "img2.png": ...
}

其中键为图片名称。

我能从list_files返回的路径中提取出图片名称,但这个名称是tf.string类型,没法直接用来访问标注字典里的值。

请问有没有简便方法把tf.string转成Python字符串,从而读取JSON里的真值数据?或者有没有办法把标注转换成适合TensorFlow使用的类型?

附上相关代码示例:

from typing import Mapping
from numpy import ndarray
import tensorflow as tf
import cv2 as cv
from pathlib import Path
from typing import Any, Mapping, NamedTuple
import json

class Point:
    x: float
    y: float

    def __init__(self, x: float, y: float):
        self.x = x
        self.y = y

class BoundingBox(NamedTuple):
    top: float
    left: float
    bottom: float
    right: float

class Annotation:
    image: tf.Tensor
    bounding_box: tf.Tensor
    is_visible: bool

    def __init__(self, image, bounding_box, is_visible):
        self.image = image
        self.bounding_box = bounding_box
        self.is_visible = is_visible

LABELS = {
    "NO_CLUB": 0,
    "CLUB": 1,
    "bbox": BoundingBox,
}


def is_in_split(image_path: tf.string, is_training: bool) -> bool:
    hash = tf.strings.to_hash_bucket_fast(image_path, 10)
    if is_training:
        return hash < 8
    else:
        return hash >= 8


def create_image_and_annotation(image_path: tf.string, annotation: Mapping[str, Any]):
    bits = tf.io.read_file(image_path)
    file_split = tf.strings.split(image_path, "/")
    image_name = file_split[-1]
    suffix = tf.strings.split(image_name, ".")[-1]

    jpeg = [
        tf.convert_to_tensor("jpg", dtype=tf.string),
        tf.convert_to_tensor("JPG", dtype=tf.string),
        tf.convert_to_tensor("jpeg", dtype=tf.string),
        tf.convert_to_tensor("JPEG", dtype=tf.string),
    ]
    is_jpeg = [tf.math.equal(suffix, s) for s in jpeg]
    png = [
        tf.convert_to_tensor("png", dtype=tf.string),
        tf.convert_to_tensor("PNG", dtype=tf.string),
    ]
    is_png = [tf.math.equal(suffix, s) for s in png]
    if tf.math.reduce_any(is_jpeg):
        image = tf.io.decode_jpeg(bits, channels=3)
    else:
        image = tf.io.decode_png(bits, channels=3)
    # 这里想用image_name获取对应图片的标注!<---
    bounding_box = BoundingBox(0,0,10,10)
    return image, (bounding_box, True)


def createDataset(dir: Path, annotation: Mapping[str, Any], is_training: bool) -> tf.data.Dataset:
    image_path_png = str(dir / "images" / "*.png")
    image_path_PNG = str(dir / "images" / "*.PNG")
    image_path_jpg = str(dir / "images" / "*.jpg")
    image_path_JPG = str(dir / "images" / "*.JPG")
    image_path_jpeg = str(dir / "images" / "*.jpeg")
    image_path_JPEG = str(dir / "images" / "*.JPEG")
    image_dirs = [image_path_png, image_path_PNG, image_path_jpg, image_path_JPG, image_path_jpeg, image_path_JPEG]

    dataset = (tf.data.Dataset.list_files(image_dirs)
        .shuffle(1000)
        .map(lambda x: create_image_and_annotation(x, annotation))
        )

    for d in dataset:
        pass

    return dataset

def getDataset(data_root_path: Path, is_training: bool) -> tf.data.Dataset:
        dirs = [x for x in data_root_path.iterdir() if x.is_dir()]
        datasets = []
        for dir in dirs:
            json_path = dir / "annotations.json"
            with open(json_path) as json_file:
                annotation = json.load(json_file)
                createDataset(dir, annotation, is_training=is_training)



training_data = getDataset(Path("/home/erik/Datasets/ClubHeadDetection"), True)
解决方案

方法一:用tf.py_function转换tf.string为Python字符串

借助tf.py_function可以在Dataset的map操作中调用普通Python函数,直接把tf.string转成Python字符串后访问标注字典。

修改create_image_and_annotation函数,把标注读取逻辑拆分到Python函数中:

def load_annotation(image_name_str: str, annotation_dict: Mapping[str, Any]):
    # 用Python字符串直接查询标注字典
    ann_data = annotation_dict[image_name_str]
    # 解析标注为BoundingBox和可见性标记
    bbox = BoundingBox(
        top=ann_data["top"],
        left=ann_data["left"],
        bottom=ann_data["bottom"],
        right=ann_data["right"]
    )
    is_visible = ann_data.get("is_visible", True)
    return bbox, is_visible

def create_image_and_annotation(image_path: tf.string, annotation: Mapping[str, Any]):
    bits = tf.io.read_file(image_path)
    file_split = tf.strings.split(image_path, "/")
    image_name = file_split[-1]
    suffix = tf.strings.split(image_name, ".")[-1]

    jpeg = [
        tf.convert_to_tensor("jpg", dtype=tf.string),
        tf.convert_to_tensor("JPG", dtype=tf.string),
        tf.convert_to_tensor("jpeg", dtype=tf.string),
        tf.convert_to_tensor("JPEG", dtype=tf.string),
    ]
    is_jpeg = [tf.math.equal(suffix, s) for s in jpeg]
    png = [
        tf.convert_to_tensor("png", dtype=tf.string),
        tf.convert_to_tensor("PNG", dtype=tf.string),
    ]
    is_png = [tf.math.equal(suffix, s) for s in png]
    if tf.math.reduce_any(is_jpeg):
        image = tf.io.decode_jpeg(bits, channels=3)
    else:
        image = tf.io.decode_png(bits, channels=3)
    
    # 用tf.py_function调用Python逻辑,转换tf.string为Python字符串
    bounding_box, is_visible = tf.py_function(
        func=lambda name: load_annotation(name.numpy().decode('utf-8'), annotation),
        inp=[image_name],
        Tout=[tf.float32, tf.bool]  # 按实际标注类型定义输出Tensor类型
    )
    # 手动设置张量形状,避免TensorFlow自动推断出错
    bounding_box.set_shape((4,))
    return image, (bounding_box, is_visible)

注意:这种方法会跳出TensorFlow计算图,可能影响性能,也无法直接导出SavedModel,适合小数据集或快速验证场景。

方法二:将标注转换为TensorFlow哈希表(tf.lookup.StaticHashTable)

把JSON标注转换成TensorFlow原生的哈希表,全程在计算图内操作,性能更高且支持模型导出。

修改getDataset和createDataset函数,提前构建哈希表:

def createDataset(dir: Path, annotation_table: tf.lookup.StaticHashTable, is_training: bool) -> tf.data.Dataset:
    image_path_png = str(dir / "images" / "*.png")
    image_path_PNG = str(dir / "images" / "*.PNG")
    image_path_jpg = str(dir / "images" / "*.jpg")
    image_path_JPG = str(dir / "images" / "*.JPG")
    image_path_jpeg = str(dir / "images" / "*.jpeg")
    image_path_JPEG = str(dir / "images" / "*.JPEG")
    image_dirs = [image_path_png, image_path_PNG, image_path_jpg, image_path_JPG, image_path_jpeg, image_path_JPEG]

    def process_image(image_path):
        bits = tf.io.read_file(image_path)
        file_split = tf.strings.split(image_path, "/")
        image_name = file_split[-1]
        suffix = tf.strings.split(image_name, ".")[-1]

        jpeg = [
            tf.convert_to_tensor("jpg", dtype=tf.string),
            tf.convert_to_tensor("JPG", dtype=tf.string),
            tf.convert_to_tensor("jpeg", dtype=tf.string),
            tf.convert_to_tensor("JPEG", dtype=tf.string),
        ]
        is_jpeg = [tf.math.equal(suffix, s) for s in jpeg]
        png = [
            tf.convert_to_tensor("png", dtype=tf.string),
            tf.convert_to_tensor("PNG", dtype=tf.string),
        ]
        is_png = [tf.math.equal(suffix, s) for s in png]
        if tf.math.reduce_any(is_jpeg):
            image = tf.io.decode_jpeg(bits, channels=3)
        else:
            image = tf.io.decode_png(bits, channels=3)
        
        # 用哈希表查询标注,返回JSON字符串后解析为Tensor
        ann_str = annotation_table.lookup(image_name)
        ann_data = tf.io.parse_json(ann_str, {
            "top": tf.float32,
            "left": tf.float32,
            "bottom": tf.float32,
            "right": tf.float32,
            "is_visible": tf.bool
        })
        bounding_box = tf.stack([ann_data["top"], ann_data["left"], ann_data["bottom"], ann_data["right"]])
        is_visible = ann_data["is_visible"]
        return image, (bounding_box, is_visible)

    dataset = (tf.data.Dataset.list_files(image_dirs)
        .shuffle(1000)
        .map(process_image, num_parallel_calls=tf.data.AUTOTUNE)
        )

    return dataset

def getDataset(data_root_path: Path, is_training: bool) -> tf.data.Dataset:
    dirs = [x for x in data_root_path.iterdir() if x.is_dir()]
    datasets = []
    for dir in dirs:
        json_path = dir / "annotations.json"
        with open(json_path) as json_file:
            annotation = json.load(json_file)
            # 转换标注为哈希表的键值对
            keys = tf.convert_to_tensor(list(annotation.keys()), dtype=tf.string)
            # 把每个标注转为JSON字符串,方便TensorFlow解析
            values = tf.convert_to_tensor([json.dumps(v) for v in annotation.values()], dtype=tf.string)
            # 创建哈希表,设置默认标注(避免找不到对应键报错)
            annotation_table = tf.lookup.StaticHashTable(
                tf.lookup.KeyValueTensorInitializer(keys, values),
                default_value=json.dumps({"top":0.0, "left":0.0, "bottom":0.0, "right":0.0, "is_visible":False})
            )
            ds = createDataset(dir, annotation_table, is_training=is_training)
            datasets.append(ds)
    # 合并所有子数据集
    return tf.data.Dataset.concatenate(*datasets)

这种方法完全基于TensorFlow计算图,没有Python开销,适合大规模训练场景,也能正常导出模型。


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.12 08:42:23