Team Ai
Datasetpublic

GD-ML/AndroidCode

sourceHugging Facemitupdated 6mo agoView on Hugging Face
1likes1.4kdownloads
androidcontrol_data_extract.py106 linesDownload Raw Back to root
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()