GD-ML/AndroidCode
11.4k
1import tensorflow as tf2import os3import glob4from tqdm import tqdm5 6# ================= 配置区域 =================7INPUT_DIR = "./downloads" # 下载数据的目录8OUTPUT_DIR = "./images" # 结果保存目录9 10# 设置为 None 表示提取所有数据11# 设置为 具体数字(如 100)用于测试12MAX_EPISODES = None 13# ===========================================14 15# 定义解析格式16feature_description = {17 'episode_id': tf.io.FixedLenFeature([], tf.int64),18 'screenshots': tf.io.VarLenFeature(tf.string),19}20 21def parse_func(example_proto):22 return tf.io.parse_single_example(example_proto, feature_description)23 24def main():25 # 屏蔽掉那些烦人的 TensorFlow 日志(只显示 Error)26 os.environ['TF_CPP_MIN_LOG_LEVEL'] = '2' 27 28 try:29 tf.config.set_visible_devices([], 'GPU')30 except:31 pass32 33 # 1. 扫描文件34 tfrecord_files = sorted(glob.glob(os.path.join(INPUT_DIR, "android_control-*-of-00020")))35 if not tfrecord_files:36 print("❌ 未找到数据文件。")37 return38 39 print(f"📂 找到 {len(tfrecord_files)} 个 TFRecord 文件。")40 print(f"📂 图片将保存至: {OUTPUT_DIR}")41 42 # 统计变量43 total_images_extracted = 044 global_episode_count = 045 46 # 初始化进度条47 # 如果 MAX_EPISODES 是 None,进度条只显示计数,不显示百分比(因为不知道总数)48 pbar = tqdm(total=MAX_EPISODES, unit="ep", desc="提取进度")49 50 # 2. 遍历所有文件51 for file_path in tfrecord_files:52 if MAX_EPISODES is not None and global_episode_count >= MAX_EPISODES:53 break54 55 # 加载数据集56 raw_dataset = tf.data.TFRecordDataset(file_path, compression_type='GZIP')57 parsed_dataset = raw_dataset.map(parse_func)58 59 # 3. 遍历每一个 Episode60 for features in parsed_dataset:61 if MAX_EPISODES is not None and global_episode_count >= MAX_EPISODES:62 break63 64 ep_id = features['episode_id'].numpy()65 66 # 创建文件夹67 ep_dir = os.path.join(OUTPUT_DIR, str(ep_id))68 69 # 断点续传检查70 if os.path.exists(ep_dir):71 # 如果跳过,是否要算在进度条里?建议算上,以免看起来像卡住了72 # 但如果不读取文件,我们不知道这个跳过的 Episode 里有几张图73 # 所以这里简单更新一下 episode 计数74 # pbar.update(1) 75 continue 76 77 os.makedirs(ep_dir, exist_ok=True)78 79 # 获取图片数据80 screenshots = features['screenshots'].values.numpy()81 82 # 4. 保存图片83 for i, img_bytes in enumerate(screenshots):84 img_name = f"{ep_id}_{i}.png"85 save_path = os.path.join(ep_dir, img_name)86 with open(save_path, "wb") as f:87 f.write(img_bytes)88 89 # --- 更新统计数据 ---90 current_imgs_count = len(screenshots)91 total_images_extracted += current_imgs_count92 global_episode_count += 193 94 # 更新进度条95 pbar.update(1)96 # 在进度条后面实时显示总图片数97 pbar.set_postfix({"Total Images": total_images_extracted})98 99 pbar.close()100 print(f"\n✅ 任务结束!")101 print(f" 共处理 Episodes: {global_episode_count}")102 print(f" 共提取 Images: {total_images_extracted}")103 104 105if __name__ == "__main__":106 main()