如何在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
相关产品推荐
相关产品推荐

