博客

  • 从docker hub镜像中拉取镜像

    最近换了Apple M4的MacBook,其他使用都还好,就是镜像拉群经常出问题,主要是docker hub的镜像有问题(忘记了旧电脑怎么配置的),要么就是报no matching manifest。

    Error response from daemon: no matching manifest for linux/arm64/v8 in the manifest list entries: no match for platform in manifest: not found

    查了下M4是ARMv9架构,具体是ARMv9.2-A,如果没有合适的镜像可以用linux/amd64试试。

    专门写了一个脚本来处理这个问题

    #!/bin/bash
    
    # Docker镜像仓库列表
    MIRRORS=(
        这里放自己常用的镜像
    )
    
    # 检查参数
    if [ $# -ne 1 ]; then
        echo "使用方法: $0 <imagename>"
        echo "示例: $0 nginx:latest"
        exit 1
    fi
    
    IMAGE_NAME="$1"
    
    echo "🚀 开始为镜像 $IMAGE_NAME 选择最快的镜像源..."
    
    # 随机选择3个镜像源
    selected_mirrors=()
    temp_mirrors=("${MIRRORS[@]}")
    
    for i in {1..3}; do
        if [ ${#temp_mirrors[@]} -eq 0 ]; then
            break
        fi
    
        # 生成随机索引
        random_index=$((RANDOM % ${#temp_mirrors[@]}))
        selected_mirrors+=("${temp_mirrors[$random_index]}")
    
        # 从临时数组中移除已选择的镜像源
        temp_mirrors=("${temp_mirrors[@]:0:$random_index}" "${temp_mirrors[@]:$((random_index + 1))}")
    done
    
    echo "📡 随机选择的镜像源: ${selected_mirrors[*]}"
    
    # 测试延迟并找到最快的镜像源
    fastest_mirror=""
    fastest_time=9999
    
    echo "🏃 正在测试镜像源延迟..."
    
    for mirror in "${selected_mirrors[@]}"; do
        echo -n "测试 $mirror ... "
    
        # ping测试,取3次平均值
        ping_result=$(ping -c 3 -W 2 "$mirror" 2>/dev/null | grep "avg" | awk -F'/' '{print $5}')
    
        if [ -n "$ping_result" ]; then
            echo "${ping_result}ms"
    
            # 使用awk进行浮点数比较
            is_faster=$(awk -v current="$ping_result" -v fastest="$fastest_time" 'BEGIN {print (current < fastest) ? 1 : 0}')
    
            if [ "$is_faster" -eq 1 ]; then
                fastest_time="$ping_result"
                fastest_mirror="$mirror"
            fi
        else
            echo "超时"
        fi
    done
    
    # 检查是否找到可用的镜像源
    if [ -z "$fastest_mirror" ]; then
        echo "❌ 所有选择的镜像源都无法访问,尝试直接拉取..."
        fastest_mirror=""
    else
        echo "🎯 最快的镜像源: $fastest_mirror (${fastest_time}ms)"
    fi
    
    # 构建完整的镜像名称
    if [ -n "$fastest_mirror" ]; then
        full_image_name="$fastest_mirror/$IMAGE_NAME"
    else
        full_image_name="$IMAGE_NAME"
    fi
    
    echo "📥 开始拉取镜像: $full_image_name"
    
    # 拉取镜像
    if docker pull --platform linux/amd64 "$full_image_name"; then
        echo "✅ 镜像拉取成功!"
    
        # 如果使用了镜像源,则重新标记为原始名称
        if [ -n "$fastest_mirror" ]; then
            echo "🏷️  重新标记镜像为: $IMAGE_NAME"
            if docker tag "$full_image_name" "$IMAGE_NAME"; then
                echo "✅ 镜像标记成功!"
    
                # 询问是否删除带镜像源前缀的镜像
                read -p "是否删除带镜像源前缀的镜像 $full_image_name? (y/N): " -n 1 -r
                echo
                if [[ $REPLY =~ ^[Yy]$ ]]; then
                    docker rmi "$full_image_name"
                    echo "🗑️  已删除镜像: $full_image_name"
                fi
            else
                echo "❌ 镜像标记失败!"
                exit 1
            fi
        fi
    
        echo "🎉 完成! 镜像 $IMAGE_NAME 已准备就绪"
    else
        echo "❌ 镜像拉取失败!"
        exit 1
    fi

    拉取效果

    运行效果

  • 用 GitHub Actions 实现 Docker 多架构镜像构建与 Manifest 合并

    用 GitHub Actions 实现 Docker 多架构镜像构建与 Manifest 合并

    最近在折腾一个开源项目的时候,碰到了Docker 镜像构建问题。现在CPU 架构越来越多样,除了我们熟知的 amd64 (或者叫 x86_64),arm64 架构也因为苹果的 M系列芯片、各种云服务器实例以及树莓派等嵌入式设备的普及而变得越来越重要。如果我们的 Docker 镜像只支持单一架构,那显然是不行的。

    最简单的方式就是在docker action中指定platform,搭配QEMU可以全自动的实现多架构镜像构建,但是QEMU很慢,会极大增加构建时间。

    本文介绍依赖原生runner实现多架构镜像构建的方式。

    整个 Workflow 主要包含两个核心的 Job:

    1. build-and-push:这个 Job 会并行地为我们指定的多个平台(例如 linux/amd64 和 linux/arm64)分别构建 Docker 镜像。构建完成后,它会将镜像推送到 GitHub Container Registry (GHCR),并把每个平台镜像的 digest(摘要)作为 artifact 上传。
    2. merge:这个 Job 会在所有平台的 build-and-push Job 成功完成后执行。它会下载之前上传的各个平台的 digest 文件,然后使用 docker buildx imagetools 命令创建一个 Manifest List。这个 Manifest List 会将不同架构的镜像关联到一个统一的镜像标签下(例如 your-repo/image:latest),用户拉取这个标签时,Docker 会自动根据当前系统架构选择合适的镜像。

    Job 1: build-and-push – 分平台构建与推送

    这个 Job 的核心在于它的 strategy.matrix 配置,它让我们能够轻松地为不同的操作系统和平台组合并行执行构建任务。

    jobs:
      build-and-push:
        strategy:
          matrix:
            include:
              - os: ubuntu-latest
                platform: linux/amd64
              - os: ubuntu-24.04-arm 
                platform: linux/arm64
        runs-on: ${{ matrix.os }}
        permissions:
          contents: read
          packages: write 

    这里定义了两个构建组合:一个在 ubuntu-latest (通常是 amd64)上构建 linux/amd64 镜像,另一个在 ubuntu-24.04-arm (GitHub 提供的 ARM runner)上构建 linux/arm64 镜像。

    构建并推送 Docker 镜像阶段使用了 docker/build-push-action@v6

          - name: Build and push Docker image
            uses: docker/build-push-action@v6
            id: build
            with:
              context: .
              platforms: ${{ matrix.platform }} # 指定当前矩阵的平台
              push: ${{ github.event_name != 'pull_request' }} # PR时不推送
              annotations: ${{ steps.meta.outputs.annotations }}
              labels: ${{ steps.meta.outputs.labels }}
              outputs: type=image,name=${{ env.GHCR_IMAGE }},push-by-digest=true,name-canonical=true,push=${{ github.event_name != 'pull_request' }},oci-mediatypes=true
              cache-from: type=gha,scope=${{ github.repository }}-${{ github.ref_name }}-${{ matrix.platform }}
              cache-to: type=gha,mode=max,scope=${{ github.repository }}-${{ github.ref_name }}-${{ matrix.platform }}
    • platforms: ${{ matrix.platform }}: 明确告诉 Buildx 为当前矩阵指定的平台构建镜像。
    • push: ${{ github.event_name != ‘pull_request’ }}: 再次确认只有非 PR 事件才推送。
    • outputs: type=image,…,push-by-digest=true,…: 这点非常关键! push-by-digest=true 确保了镜像是基于其内容摘要 (digest) 推送的,而不是基于可变的标签。每个平台构建的镜像都会有一个唯一的、以 sha256:… 开头的 digest。这是后续合并 Manifest 的基础,因为 Manifest List 就是通过这些 digest 来引用不同平台的具体镜像。
    • cache-from 和 cache-to: 开启了 Docker 构建缓存,利用 GitHub Actions 的缓存机制来加速后续的构建,非常实用。

    构建并推送完成后,steps.build.outputs.digest 会输出镜像的 digest。我们将这个 digest (去掉了 sha256: 前缀) 作为文件名,创建一个空文件,然后通过 actions/upload-artifact@v4 将其作为 artifact 上传。文件名中包含了平台信息 (${{ env.PLATFORM_PAIR }}),方便后续 Job 区分。保留时间设置为1天,足够后续 merge Job 使用了。

      - name: Export digest
        run: |
          mkdir -p /tmp/digests
          digest="${{ steps.build.outputs.digest }}"
          touch "/tmp/digests/${digest#sha256:}" # 将sha256:前缀去掉作为文件名
    
      - name: Upload artifact
        uses: actions/upload-artifact@v4
        with:
          name: digests-${{ env.PLATFORM_PAIR }}
          path: /tmp/digests/*
          if-no-files-found: error
          retention-days: 1

    Job 2: merge – 合并 Manifest List

    使用 actions/download-artifact@v4 下载之前所有平台上传的 digest 文件。这些文件会被放到 /tmp/digests 目录下。

          - name: Download digests
            uses: actions/download-artifact@v4
            with:
              path: /tmp/digests # 下载到指定路径
              pattern: digests-* # 匹配之前上传的 artifact 名称
              merge-multiple: true # 如果有多个同名 artifact (理论上这里不会),会合并

    然后再次使用 docker/metadata-action@v5,但这次的目的是为最终的 Manifest List 生成合适的标签。

            id: meta
            uses: docker/metadata-action@v5
            with:
              images: ${{ env.GHCR_IMAGE }}
              annotations: | # OCI annotations for the manifest list itself
                type=org.opencontainers.image.description,value=${{ github.event.repository.description || 'No description provided' }}
              tags: | # 多种打标策略
                type=semver,pattern={{version}} # v1.2.3
                type=semver,pattern={{major}}.{{minor}} # v1.2
                type=sha,format=short # 短sha
                type=ref,event=branch # 分支名 (e.g., main)
                latest # 始终打 latest 标签

    最后合并推送

          - name: Create manifest list and pushs
            working-directory: /tmp/digests # 切换到 digests 存放的目录
            id: manifest-annotate
            continue-on-error: true # 如果带注解的失败了,允许继续
            run: |
                  docker buildx imagetools create \
                    $(jq -cr '.tags | map("-t " + .) | join(" ")' <<< "$DOCKER_METADATA_OUTPUT_JSON") \ # 从 metadata action 的输出中提取所有 tags
                    --annotation='index:org.opencontainers.image.description=${{ github.event.repository.description }}' \
                    --annotation='index:org.opencontainers.image.created=${{ steps.timestamp.outputs.timestamp }}' \
                    --annotation='index:org.opencontainers.image.url=${{ github.event.repository.url }}' \
                    --annotation='index:org.opencontainers.image.source=${{ github.event.repository.url }}' \
                    $(printf '${{ env.GHCR_IMAGE }}@sha256:%s ' *) # 将 /tmp/digests 下的所有文件名(即 digests)拼接到命令中

    这种方案的提升是巨大的,比如同样的项目,使用Qemu方案需要12分钟,切换到arm runner只需要4分钟。

  • 使用tensorflow检测睡眠鼻鼾情况

    使用tensorflow检测睡眠鼻鼾情况

    最近需要监测下睡眠情况,主要是分析打呼噜的情况。家里有一个小米摄像头,正好利用起来。

    步骤也比较简单,睡觉前摄像头打开,然后随便对着墙(因为我们只要音频),第二天起床后把所有监控文件按照时间顺序合并,并转为wav文件。

    我这里7个小时左右的音频,大小2.4GB。

    由于分析过程中发现有咳嗽的情况,又增加了咳嗽的监测。

    本来想一次性分析的,结果OOM了(我用的虚拟机,分配了48G),只有改成分段处理,内存消耗大概4GB。

    完整代码

    import tensorflow as tf
    import tensorflow_hub as hub
    import numpy as np
    import librosa
    import pandas as pd
    import os
    # import soundfile as sf # librosa.load is generally robust enough
    
    # 可调整参数:
    # audio_file_name: 在 main() 函数中修改为您的音频文件名。
    # event_confidence_thresholds: 在 main() 函数中调整。这是一个字典,键是事件标签 (例如 "Snoring", "Cough"),
    #                             值是介于 0 和 1 之间的置信度阈值。较高的值会使得检测更严格,
    #                             减少误报,但可能漏掉不典型的或轻微的声音。
    #                             建议从 Snoring: 0.15-0.25, Cough: 0.2-0.3 开始尝试。
    # event_merge_threshold: 在 main() 函数中调整。如果同一类型的两个检测事件的结束和开始时间间隔
    #                       小于此阈值(秒),它们将被合并。
    # CHUNK_DURATION_SEC: 在 predict_audio_events 中调整,处理长音频时的分块大小(秒)。
    
    def load_yamnet_model_and_class_names():
        """加载 YAMNet 模型和类别名称"""
        try:
            print("正在加载 YAMNet 模型...")
            yamnet_model_handle = 'https://tfhub.dev/google/yamnet/1'
            yamnet_model = hub.load(yamnet_model_handle)
            print("YAMNet 模型加载成功。")
    
            print("正在加载 YAMNet 类别名称...")
            class_map_path = yamnet_model.class_map_path().numpy().decode('utf-8')
            class_names_df = pd.read_csv(class_map_path)
            class_names = class_names_df['display_name'].tolist()
            print(f"YAMNet 类别名称加载成功 ({len(class_names)} 个类别)。")
            return yamnet_model, class_names
        except Exception as e:
            print(f"加载 YAMNet 模型或类别名称时出错: {e}")
            print("请确保您的网络连接正常,并且 TensorFlow Hub 可以访问。")
            print("如果问题持续,您可能需要检查 TensorFlow 和 TensorFlow Hub 的版本兼容性。")
            return None, None
    
    def merge_overlapping_events(events, time_threshold=0.5):
        """
        合并时间上重叠或非常接近的事件。
        Args:
            events (list): 事件字典列表,每个字典包含 "start_time_seconds",
                           "end_time_seconds", "confidence", "label"。
            time_threshold (float): 合并事件的最大时间间隔(秒)。
        Returns:
            list: 合并后的事件列表。
        """
        if not events:
            return []
    
        events.sort(key=lambda x: x["start_time_seconds"])
        merged = []
        current_event = events[0].copy()
    
        for i in range(1, len(events)):
            next_event = events[i]
            if next_event["start_time_seconds"] <= current_event["end_time_seconds"] + time_threshold:
                current_event["end_time_seconds"] = max(current_event["end_time_seconds"], next_event["end_time_seconds"])
                current_event["confidence"] = max(current_event["confidence"], next_event["confidence"])
            else:
                merged.append(current_event)
                current_event = next_event.copy()
        merged.append(current_event)
        return merged
    
    def predict_audio_events(audio_path, model, class_names_list,
                             target_labels_with_thresholds,
                             merge_time_threshold=0.5,
                             default_confidence_threshold=0.1,
                             chunk_duration_sec=60): # <<< 新增:分块处理时长(秒)
        """
        使用 YAMNet 检测音频文件中指定类型的声音事件,支持长音频分块处理。
        """
        if not model or not class_names_list:
            print("错误:YAMNet 模型或类别列表未加载。")
            return {label: [] for label in target_labels_with_thresholds}
    
        if not os.path.exists(audio_path):
            print(f"错误:音频文件未找到: {audio_path}")
            return {label: [] for label in target_labels_with_thresholds}
    
        label_indices = {}
        for target_label in target_labels_with_thresholds.keys():
            try:
                label_indices[target_label] = class_names_list.index(target_label)
            except ValueError:
                print(f"警告:标签 '{target_label}' 在 YAMNet 类别名称中未找到。将跳过此标签。")
        
        if not label_indices:
            print("错误:没有可用的有效目标标签。")
            return {label: [] for label in target_labels_with_thresholds}
    
        print(f"正在加载音频文件: {audio_path}...")
        try:
            # YAMNet 需要 16kHz 单声道,float32 范围 [-1.0, 1.0]
            # librosa.load 会自动重采样到16kHz
            full_waveform, sr = librosa.load(audio_path, sr=16000, mono=True)
            full_waveform = full_waveform.astype(np.float32)
            if np.max(np.abs(full_waveform)) > 1.0:
                 full_waveform /= np.max(np.abs(full_waveform))
        except Exception as e:
            print(f"加载或转换音频文件时出错: {e}")
            return {label: [] for label in target_labels_with_thresholds}
    
        if full_waveform.size == 0:
            print("错误:加载的音频波形为空。")
            return {label: [] for label in target_labels_with_thresholds}
        
        print(f"音频加载完毕。采样率: {sr}Hz, 总时长: {len(full_waveform)/sr:.2f} 秒。")
    
        frame_hop_seconds = 0.48  # YAMNet 帧移
        frame_window_seconds = 0.96 # YAMNet 窗长
    
        # 分块处理
        samples_per_chunk = int(chunk_duration_sec * sr)
        num_samples_total = len(full_waveform)
        
        all_scores_list = []
        
        print(f"总样本数: {num_samples_total}. 每块样本数: {samples_per_chunk} (对应 {chunk_duration_sec} 秒)")
        print("开始使用 YAMNet 进行分块预测...")
    
        for i in range(0, num_samples_total, samples_per_chunk):
            chunk_start_sample = i
            chunk_end_sample = min(i + samples_per_chunk, num_samples_total)
            chunk_waveform = full_waveform[chunk_start_sample:chunk_end_sample]
    
            if len(chunk_waveform) < int(frame_window_seconds * sr) : # YAMNet需要至少一个完整窗口的音频
                 if num_samples_total < int(frame_window_seconds * sr) and i == 0: # 音频本身就太短
                     print(f"音频片段 (从样本 {chunk_start_sample}) 太短 ({len(chunk_waveform)/sr:.2f}s),无法处理。至少需要 {frame_window_seconds:.2f}s。")
                 elif len(chunk_waveform) > 0 : # 最后一个块可能很短,但仍尝试处理
                     print(f"处理最后一个短音频片段 (从样本 {chunk_start_sample}, 时长 {len(chunk_waveform)/sr:.2f}s)...")
                 else: # 空块,跳过
                     continue
            else:
                print(f"处理块: 样本 {chunk_start_sample} 到 {chunk_end_sample} (时长 {len(chunk_waveform)/sr:.2f}s)")
    
            if len(chunk_waveform) > 0:
                # YAMNet 模型直接处理波形
                # scores的形状是 (num_frames, num_classes)
                scores_chunk, _, _ = model(chunk_waveform)
                all_scores_list.append(scores_chunk.numpy())
            else:
                print(f"跳过空音频块 (样本 {chunk_start_sample} 到 {chunk_end_sample})")
                
        if not all_scores_list:
            print("警告:没有生成任何分数。音频可能太短或在分块后为空。")
            return {label: [] for label in label_indices.keys()}
    
        print("所有块处理完毕,正在合并分数...")
        scores_np = np.concatenate(all_scores_list, axis=0)
        print(f"合并后的总帧数: {scores_np.shape[0]}")
    
    
        all_detected_events = {label: [] for label in label_indices.keys()}
    
        print(f"根据合并后的分数处理 {scores_np.shape[0]} 个音频帧...")
        for i, frame_scores in enumerate(scores_np):
            for target_label, class_index in label_indices.items():
                score = frame_scores[class_index]
                confidence_threshold = target_labels_with_thresholds.get(target_label, default_confidence_threshold)
                
                if score >= confidence_threshold:
                    # 时间戳是相对于整个音频的开始
                    start_time = i * frame_hop_seconds
                    end_time = start_time + frame_window_seconds
                    all_detected_events[target_label].append({
                        "start_time_seconds": round(start_time, 2),
                        "end_time_seconds": round(end_time, 2),
                        "confidence": round(float(score), 3),
                        "label": target_label
                    })
    
        merged_results = {}
        for target_label, events in all_detected_events.items():
            if events:
                print(f"初步检测到 {len(events)} 个可能的 '{target_label}' 事件(合并前)。")
                print(f"正在为 '{target_label}' 合并时间间隔小于 {merge_time_threshold} 秒的事件...")
                merged_events = merge_overlapping_events(events, time_threshold=merge_time_threshold)
                print(f"合并后得到 {len(merged_events)} 个 '{target_label}' 事件。")
                merged_results[target_label] = merged_events
            else:
                merged_results[target_label] = []
                print(f"未初步检测到 '{target_label}' 事件。")
    
        return merged_results
    
    def main():
        audio_file_name = "sleep.wav" # <<--- 将此替换为您的文件名
    
        yamnet_model, class_names = load_yamnet_model_and_class_names()
    
        if yamnet_model and class_names:
            print(f"\n开始分析音频文件: {audio_file_name}")
    
            target_events_with_thresholds = {
                "Snoring": 0.1,
                "Cough": 0.25
            }
            event_merge_threshold = 1.0
            audio_processing_chunk_seconds = 600 # 处理音频的块大小,单位秒
    
            detected_events_map = predict_audio_events(
                audio_file_name,
                yamnet_model,
                class_names,
                target_labels_with_thresholds=target_events_with_thresholds,
                merge_time_threshold=event_merge_threshold,
                default_confidence_threshold=0.15,
                chunk_duration_sec=audio_processing_chunk_seconds # 传递块大小
            )
    
            print("\n--- 事件检测结果 ---")
            any_event_detected = False
            for label, events in detected_events_map.items():
                if events:
                    any_event_detected = True
                    print(f"\n--- 检测到 '{label}' 事件 ---")
                    total_event_duration = 0
                    for event in events:
                        duration = event['end_time_seconds'] - event['start_time_seconds']
                        total_event_duration += duration
                        print(
                            f"标签: {event['label']}, 从 {event['start_time_seconds']:.2f} 秒 "
                            f"到 {event['end_time_seconds']:.2f} 秒 "
                            f"(时长: {duration:.2f} 秒), "
                            f"最大置信度: {event['confidence']:.3f}"
                        )
                    print(f"\n'{label}' 事件总时长 (近似): {total_event_duration:.2f} 秒")
                    print(f"共检测到 {len(events)} 个 '{label}' 片段。")
                else:
                    confidence_val = target_events_with_thresholds.get(label)
                    if confidence_val is None: # Should not happen if label is in target_events_with_thresholds
                        confidence_val_str = f"默认({0.15})" # Assuming 0.15 is the default_confidence_threshold
                    else:
                        confidence_val_str = str(confidence_val)
                    print(f"\n在文件 '{audio_file_name}' 中未检测到明显的 '{label}' 事件(使用阈值 {confidence_val_str})。")
            
            if not any_event_detected:
                 print(f"\n在文件 '{audio_file_name}' 中未检测到任何指定的目标事件。")
                 print("您可以尝试调整 `target_events_with_thresholds` 中的阈值以检测更细微的声音,但这可能会增加误报。")
    
        else:
            print("由于模型或类别名称加载失败,无法进行事件检测。")
    
    if __name__ == "__main__":
        main()

    检测效果

  • Rust将PDF转为图片

    最近有几张电子发票要报保险,但是腾讯微保上传发票需要上传图片,想着直接转一下,然后手机上各种APP试了一圈要么要收费,要么只能免费转第一页,不巧的是我这几张发票有明细表,都是两页的。网页版本有各种免费的,但是始终比较担心安全性,无奈只有自己搞一下。

    PDF的标准很复杂,自己实现显然不是最佳选择。PDFium 是一个开源的 PDF 渲染引擎,由 Google 开发和维护。它用于解析和渲染 PDF 文档,广泛应用于 Chrome 浏览器和其他项目。PDFium 提供高效的 PDF 处理功能,包括文本提取、注释、表单填充和页面渲染,支持多平台,当然就包括了Android。

    开源社区有Rust的绑定,我们可以直接使用,我这里版本用的6666,社区也有预构建文件可以直接下载。

    这里参考官方例子,唯一的区别的改动是从单页导出改为全部页面导出

    pub struct PdfToImageResult {
        pub image_path: String,
    }
    
    pub fn export_pdf_to_jpegs(path: String, out_dir: String) -> PdfToImageResult {
        let p = Pdfium::default();
    
        let document = p.load_pdf_from_file(&path, Option::None).unwrap();
    
        let render_config = PdfRenderConfig::new()
            .set_target_width(2000)
            .set_maximum_height(2000)
            .rotate_if_landscape(PdfPageRenderRotation::Degrees90, true);
    
        let mut images: Vec<RgbaImage> = Vec::new();
        let mut total_width = 0;
        let mut total_height = 0;
    
        for page in document.pages().iter() {
            let image = page
                .render_with_config(&render_config)
                .unwrap()
                .as_image()
                .to_rgba8();
            total_width = total_width.max(image.width());
            total_height += image.height();
            images.push(image);
        }
    
        let mut combined_image = RgbaImage::new(total_width, total_height);
        let mut y_offset = 0;
    
        for image in images {
            combined_image.copy_from(&image, 0, y_offset).unwrap();
            y_offset += image.height();
        }
    
        let output_path = Path::new(&out_dir)
            .join(Path::new(&path).file_stem().unwrap())
            .with_extension("png");
    
        let file = File::create(&output_path).unwrap();
        let mut writer = BufWriter::new(file);
        DynamicImage::ImageRgba8(combined_image)
            .write_to(&mut writer, ImageFormat::Png)
            .unwrap();
    
        return PdfToImageResult {
            image_path: output_path.to_str().unwrap().to_string(),
        };
    }

    链接

    https://github.com/bblanchon/pdfium-binaries

  • 使用Github Action自动发布Jellyfin插件

    Jellyfin自定义插件需要一个meta.json,内容大致如下

    {
        "category": "Metadata",
        "guid": "a3a07da4-ae5a-4d4a-a843-5aa7e3ba0a62",
        "name": "HappyMovie",
        "description": "Get metadata from tmdb.",
        "owner": "htynkn",
        "overview": "Get metadata from tmdb.",
        "targetAbi": "10.9.0.0",
        "timestamp": "2024-10-10T13:46:00Z",
        "version": "1.0.1.4"
    }

    其中targetAbi、timestamp、version字段是每个版本都不同的,其他部分可以写死。由于这些信息都可以从C#项目中获取,这里用py脚本来获取。

    tree = ET.parse("Jellyfin.Plugin.HappyMovie/Jellyfin.Plugin.HappyMovie.csproj")
    version = tree.find("./PropertyGroup/AssemblyVersion").text
    targetAbi = tree.find("./ItemGroup/*[@Include='Jellyfin.Model']").attrib["Version"]
    timestamp = datetime.now().strftime("%Y-%m-%dT%H:%M:%SZ")
    
    meta = {
        "category": "Metadata",
        "guid": "a3a07da4-ae5a-4d4a-a843-5aa7e3ba0a62",
        "name": "HappyMovie",
        "description": "Get metadata from tmdb.",
        "owner": "htynkn",
        "overview": "Get metadata from tmdb.",
        "targetAbi": f"{targetAbi}.0",
        "timestamp": timestamp,
        "version": version
    }

    另外由于插件有外部依赖,需要拷贝dll到最终产出物中,脚本添加相关拷贝最后打包

    subprocess.run([
        "dotnet",
        "build",
        "Jellyfin.Plugin.HappyMovie/Jellyfin.Plugin.HappyMovie.csproj",
        "--configuration",
        "Release"
    ])
    
    shutil.copy("Jellyfin.Plugin.HappyMovie/bin/Release/net8.0/Jellyfin.Plugin.HappyMovie.dll", f"release/{version}/")
    shutil.copy(f"{Path.home()}/.nuget/packages/yove.proxy/1.1.1/lib/netstandard2.0/Yove.Proxy.dll", f"release/{version}/")
    shutil.copy(f"{Path.home()}/.nuget/packages/tmdblib/2.2.0/lib/netstandard2.0/TMDbLib.dll", f"release/{version}/")
    
    shutil.make_archive(f"release/happymovie_{version}", "zip", f"release/{version}/")

    一般发布新版本的时候我们会打一个tag,可以基于这个特性驱动Github Action

    name: release
    
    on:
      push:
        tags:
          - '*'
    
    jobs:
      build:
        runs-on: ubuntu-latest
    
        steps:
          - name: Checkout repository
            uses: actions/checkout@v2
          - name: Setup .NET
            uses: actions/setup-dotnet@v1
            with:
              dotnet-version: 8.0.x
          - name: Set up Python 3
            uses: actions/setup-python@v2
            with:
              python-version: '3.x'
    
          - name: Execute build script
            run: |
              chmod +x ./package.py
              ./package.py
    
          - name: Create Release
            id: create_release
            uses: actions/create-release@v1
            env:
              GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }}
            with:
              tag_name: ${{ github.ref }}
              release_name: Release ${{ github.ref_name }}
              draft: true
              prerelease: false
    
          - name: Upload Release Asset
            uses: actions/upload-release-asset@v1
            env:
              GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }}
            with:
              upload_url: ${{ steps.create_release.outputs.upload_url }}
              asset_path: ./release/happymovie_${{ github.ref_name }}.zip
              asset_name: happymovie_${{ github.ref_name }}.zip
              asset_content_type: application/zip
  • Rust使用libuv库

    在使用Rust或多或少有需要调用外部库的需求,Rust FFI对于C和C++的处理有一些区别,为了简单快捷可以使用第三方工具来进行跨语言调用,比如autocxx。这里用libuv做一个演示。

    libuv 是一个跨平台的异步 I/O 库,它被设计用来作为 Node.js 的新平台抽象层。libuv 提供了跨所有主要平台的统一的非阻塞 I/O 基础设施,包括 Windows、Linux、macOS 和其他类 Unix 系统。

    我们先在github现在需要版本的libuv源代码,这里核心需要的是h文件,放到项目中,我这里用的third/libuv-1.48.0/include

    在Cargo文件中配置依赖

    [dependencies]
    autocxx = "0.26.0"
    cxx = "1.0"
    
    [build-dependencies]
    autocxx-build = "0.26.0"
    miette = { version = "5", features = ["fancy"] } 

    再新增一个build.rs文件

    fn main() -> miette::Result<()> {
        let path = std::path::PathBuf::from("src");
        let libuv_path = std::path::PathBuf::from("third/libuv-1.48.0/include");
        let mut b = autocxx_build::Builder::new("src/main.rs", &[&path, &libuv_path]).build()?;
        b.flag_if_supported("-std=c++14")
            .compile("rc-sys");
        println!("cargo:rerun-if-changed=src/main.rs");
        Ok(())
    }

    这里先尝试一个简单的,我们打印下版本号

    use autocxx::prelude::*;
    
    include_cpp! {
        #include "uv.h"
        safety!(unsafe_ffi)
        generate!("uv_default_loop")
    }
    
    fn main() {
        unsafe {
            println!("Hello, libuv! version:{}", ffi::UV_VERSION_MAJOR);
            
        }
    }
    

    直接cargo run就能看到结果。

    我们正常使用需要的是调用libuv对外暴露的方法,所以还是需要连接到libuv库。我们这里使用系统自带的,修改build.rs

    fn main() -> miette::Result<()> {
        let path = std::path::PathBuf::from("src");
        let libuv_path = std::path::PathBuf::from("third/libuv-1.48.0/include");
        let mut b = autocxx_build::Builder::new("src/main.rs", &[&path, &libuv_path]).build()?;
        b.flag_if_supported("-std=c++14")
            .compile("rc-sys");
        println!("cargo:rerun-if-changed=src/main.rs");
        println!("cargo:rustc-link-lib=uv");
        Ok(())
    }

    在rust中我们起一个loop

    use autocxx::prelude::*;
    
    include_cpp! {
        #include "uv.h"
        safety!(unsafe_ffi)
        generate!("uv_default_loop")
        generate!("uv_loop_init")
        generate!("uv_run")
        generate!("uv_loop_close")
        generate!("uv_run_mode")
    }
    
    fn main() {
        unsafe {
            println!("Hello, libuv!");
            let uv_loop_new = ffi::uv_default_loop();
            println!("Create a new loop: {:?}", uv_loop_new);
            ffi::uv_loop_init(uv_loop_new);
            println!("Initialize the loop: {:?}", uv_loop_new);
            ffi::uv_run(uv_loop_new, ffi::uv_run_mode::UV_RUN_DEFAULT);
            ffi::uv_loop_close(uv_loop_new);
            println!("Close the loop: {:?}", uv_loop_new);
            
        }
    }

    如果你更擅长C/C++,其实还可以在二次封装下libuv方法再在Rust中使用,可以避免很多问题。我们新增一个extras.h文件,并添加相关代码

    #pragma once
    
    #include <stdio.h>
    #include <stdlib.h>
    #include "uv.h"
    
    
    inline int r_start_tcp_server() {
      uv_loop_t *loop;
    
      loop = uv_default_loop();
    
      uv_tcp_t server;
      uv_tcp_init(loop, &server);
    
      struct sockaddr_in bind_addr;
      uv_ip4_addr("0.0.0.0", 7000, &bind_addr);
      
      uv_tcp_bind(&server, (const struct sockaddr *)&bind_addr, 0);
      int r = uv_listen((uv_stream_t*) &server, 128, NULL);
      if (r) {
          fprintf(stderr, "Listen error %s\n", uv_err_name(r));
          return 1;
      }
      return uv_run(loop, UV_RUN_DEFAULT);
    }

    在build.rs中添加extras.h的信息

    fn main() -> miette::Result<()> {
        let path = std::path::PathBuf::from("src");
        let libuv_path = std::path::PathBuf::from("third/libuv-1.48.0/include");
        let mut b = autocxx_build::Builder::new("src/main.rs", &[&path, &libuv_path]).build()?;
        b.flag_if_supported("-std=c++14")
            .compile("rc-sys");
        println!("cargo:rerun-if-changed=src/main.rs");
        println!("cargo:rerun-if-changed=src/extras.h");
        println!("cargo:rustc-link-lib=uv");
        Ok(())
    }

    在Rust中直接启动

    use autocxx::prelude::*;
    
    include_cpp! {
        #include "uv.h"
        #include "extras.h"
        safety!(unsafe_ffi)
        generate!("r_start_tcp_server")
    }
    
    fn main() {
        println!("Hello, libuv!");
        let code: autocxx::c_int = ffi::r_start_tcp_server();
        println!("Start_tcp_server returned: {}", code.0);
    }

    如果我们还需要编译生成libuv的话,可以使用cmake-rs,在build.rs中添加

    let dst = cmake::build("third/libuv-1.48.0");
    
    println!("cargo:rustc-link-search=native={}", dst.display());

    如果你对于cargo没有特别偏好,要使用bazel也行

    说明:这里的代码只是一个演示,现实中应该不会存在在Rust中走libuv启动TCP的情况。

  • OpenRewrite复合配方用于自动化迁移

    OpenRewrite复合配方用于自动化迁移

    OpenRewrite是一个源代码的自动重构生态系统,使开发人员能够有效地消除代码中的技术债务。

    OpenRewrite可以实现代码解析和改写,理论上可以适用于很多场景,不仅仅是消除基础债务,只要能够通过固定规则进行代码变化的工作它理论上都能胜任。本文演示用OpenRewrite将已有Spring MVC项目迁移到阿里云云原生网关中。

    背景

    当前我们有若干个微服务工程,由于早期没有网关基建,所有建设了若干个Spring MVC项目对外提供HTTP能力,这些项目运行于Jetty环境,内部通过Dubbo协议调用后端一个或者多个后端服务,由于早期研发规范缺乏,这些HTTP项目包含大量业务逻辑和编排,所以HTTP项目中的逻辑需要保留。

    从微服务长期治理看,我们希望这些逻辑更加内聚,将HTTP项目和后端的微服务合并成一个服务,通过云原生网关的能力实现HTTP和Dubbo项目的转换。

    思路

    如果是我们手动进行相关工作,大体有以下几步骤

    • 找到要迁移的Controller(一般由@Controller或者@RestController作为注解)
    • 找到要迁移的方法(一般由@RequestMapping或者@GetMapping等作为注解)
    • 提取其中的关键信息,包括请求路径和参数
    • 生成Dubbo接口定义和请求参数
    • 拷贝原有Controller方法,尽可能保留原有逻辑,只改写参数相关的逻辑
    • 拷贝原有工程代码到服务化工程中(单独的package包),并处理所有import
    • 在云原生网关配置接口信息,包括HTTP参数和后端Dubbo服务信息

    要自动化这些步骤我们需要以下技术/工具

    无损语义树

    要完成以上工作,我们需要和Java代码打交道,常规的AST解析会丢失一些信息,我们需要保留他们,包括但不限于注释、空格等,这里就需要用到OpenRewrite的无损语义树了。

    OpenRewrite 的 “Lossless Semantic Trees” 是指 OpenRewrite 使用的一种特殊的抽象语法树(AST),它能够在进行代码分析和重构时保留所有源代码的信息,包括注释、格式和空白字符。这种 AST 的设计允许 OpenRewrite 在不丢失任何原始代码细节的情况下,进行精确的代码修改。在传统的 AST 中,通常会忽略空格、换行和注释等信息,因为这些元素对于代码的语义分析不是必需的。

    OpenRewrite 的 Lossless Semantic Trees 保留了这些信息,使得代码重构操作可以像编辑器中手动修改代码一样,保持代码的原始风格和注释。这样,当使用 OpenRewrite 进行代码重构时,不仅能够保持代码逻辑的正确性,还能够保留代码的原始风格和意图,从而使重构后的代码更易于阅读和维护。

    要手动获得一个java文件的无损语义树需要手动转换一下

    ExecutionContext executionContext = initExecutionContext();
    
    JavaTypeCache javaTypeCache = new JavaTypeCache();
    JavaParser javaParser = JavaParser.fromJavaVersion()
       .classpath(CLASS_PATH_LIST).typeCache(javaTypeCache)
       .logCompilationWarningsAndErrors(false).build();
    
    javaParser.parse(javaFilePath, null, executionContext);

    自定义OpenRewrite配方

    OpenRewrite的运行关键是它的配方,内置的配方和三方开源社区的配方都无法直接满足我们的需求,所以我们需要一个复合配方。

    复合配方由多个配方组成,它们串行执行,配方直接可以通过上下文共享信息。下面是一个典型的复合配方

    package org.example.testing;
    
    import org.openrewrite.java.ChangeType;
    
    public class JUnit5Migration extends Recipe {
        @Override
        public List<Recipe> getRecipeList() {
            return Arrays.asList(
                new ChangeType("org.junit.Test", "org.junit.jupiter.api.Test", false),
                new AssertToAssertions(),
                new RemovePublicTestModifiers()
            );
        }
    }

    复合配方的运行遵循以下流程

    简而言之配方是针对每个源文件顺次运行的,如果某一步操作需要在扫描整个源代码以后再执行,则需要使用scan配方。

    配方示例

    为了更好的维护配方,我们可以把配方拆解成几个独立配方,可以参考最开始的手动迁移步骤。我们先来看一个最简单的,那就是找到我们需要的Controller

    @Value
    @EqualsAndHashCode(callSuper = false)
    public class FindTargetControllerRecipe extends ScanningRecipe<Void> {
        @Option(displayName = "Target Controller Name")
        List<String> targetControllerList;
    
        private final List<String> controllerAnnotationList = Lists.newArrayList("org.springframework.stereotype.Controller",
            "org.springframework.web.bind.annotation.RestController");
    
    
        @Override
        public TreeVisitor<?, ExecutionContext> getScanner(Void acc) {
            SharedTaskInfo sharedTaskInfo = new SharedTaskInfo();
            return new JavaIsoVisitor<ExecutionContext>() {
                @Override
                public J.ClassDeclaration visitClassDeclaration(J.ClassDeclaration classDecl, ExecutionContext executionContext) {
                    J.ClassDeclaration classDeclaration = super.visitClassDeclaration(classDecl, executionContext);
    
                    if (classDecl.getLeadingAnnotations().stream().anyMatch(anno -> {
                        return controllerAnnotationList.stream().anyMatch(annoType ->
                            TypeUtils.isOfClassType(anno.getType(), annoType));
                    })) {
                        if (targetControllerList.contains(classDecl.getSimpleName())) {
                            J.CompilationUnit cu = getCursor().getParent().getValue();
                            sharedTaskInfo.addController(ControllerInfo.newByFileAndClassDelId(cu.getId(), classDecl.getId()));
                            executionContext.putMessage(SharedTaskInfo.KEY, sharedTaskInfo);
                        }
                    }
                    return classDeclaration;
                }
            };
        }
    }

    找到以后我们就可以分析了

    @Slf4j
    public class ApiAnalysisRecipe extends ScanningRecipe<Boolean> {
       
        @Override
        public TreeVisitor<?, ExecutionContext> getScanner(Boolean acc) {
            return new JavaIsoVisitor<ExecutionContext>() {
    
                @Override
                public J.MethodDeclaration visitMethodDeclaration(J.MethodDeclaration method, ExecutionContext executionContext) {
                    J.MethodDeclaration methodDeclaration = super.visitMethodDeclaration(method, executionContext);
                    J.ClassDeclaration belongClassDecl = getCursor().firstEnclosing(J.ClassDeclaration.class);
                    SharedTaskInfo sharedTaskInfo = executionContext.getMessage(SharedTaskInfo.KEY, new SharedTaskInfo());
    
                    //获取信息存储到sharedTaskInfo
    
                    return methodDeclaration;
                }
    
                @Override
                public J.ClassDeclaration visitClassDeclaration(J.ClassDeclaration classDecl, ExecutionContext executionContext) {
                    J.ClassDeclaration classDeclaration = super.visitClassDeclaration(classDecl, executionContext);
    
                    SharedTaskInfo sharedTaskInfo = SharedTaskInfo.getFromContext(executionContext);
                    
                    //获取信息存储到sharedTaskInfo
    
                    return classDeclaration;
                }
            };
        }
    }

    然后生成一些必要的参数类,比如原始请求如下

    @Controller("/admin")
    public class HelloController {
      @RequestMapping(value = "/addUser")
      public String addUser(String name, int age) {
         //省略
      }
    }

    改写成Dubbo服务以后

    @DubboService
    public class HelloServiceImpl implments HelloService {
      @Override
      public String addUser(HelloServiceAddUserParams requestParams) {
      	String name = requestParams.getName();
      	int age = requestParams.getAge();
         //省略
      }
    }

    这里涉及到了一些代码生成,可以直接使用java parser来生成。由于配方设计到生成,需要复写generate方法

    public class GenerateWebParamRecipe extends ScanningRecipe<GenerateWebParamRecipe.Scanned> {
    
        @Override
        public Collection<? extends SourceFile> generate(GenerateWebParamRecipe.Scanned acc, ExecutionContext ctx) {
            List<SourceFile> sourceFiles = Lists.newArrayList();
    
            SharedTaskInfo sharedTaskInfo = SharedTaskInfo.getFromContext(ctx);
    
            for (ControllerInfo controllerInfo : sharedTaskInfo.getControllerInfos()) {
                for (MethodInfo methodInfo : controllerInfo.getMethodInfoMap().values()) {
                    //生成SourceFile
                }
            }
            return sourceFiles;
        }
    }

    最后复合配方如下

    
    @Value
    @EqualsAndHashCode(callSuper = false)
    public class AliyuApinDubboAllInOneRecipe extends Recipe {
        @Option
        String webPackage;
        @Option
        String newWebPackage;
        @Option
        String clientWebPackage;
        @Option
        String webProjectPath;
        @Option
        String newWebProjectPath;
        @Option
        String clientProjectPath;
        @Option
        List<String> targetControllerList;
    
        @Override
        public List<Recipe> getRecipeList() {
            List<Recipe> recipes = Lists.newLinkedList();
    
            recipes.add(new FindTargetControllerRecipe(targetControllerList));
            recipes.add(new ApiAnalysisRecipe());
            recipes.add(new GenerateDubboParamRecipe(clientWebPackage, clientProjectPath));
            recipes.add(new GenerateDuuboClientRecipe(clientWebPackage, clientProjectPath));
            recipes.add(new GenerateDubboClientImplRecipe(newWebPackage));
    
            recipes.add(new ChangePackage(webPackage, newWebPackage, true));
            recipes.add(new MoveFileRecipe(webProjectPath, newWebProjectPath));
    
            recipes.add(new SmartMavenRecipe(newWebProjectPath));
            recipes.add(new SyncApiDefineToAliyun());
    
            return recipes;
        }
    
        @Override
        public String getDisplayName() {
            return "自动迁移SpringMvc到阿里云微服务云网关";
        }
    
        @Override
        public String getDescription() {
            return "自动迁移SpringMvc到阿里云微服务云网关,包括生成接口文档、生成Dubbo客户端、生成Dubbo服务端、修改包名、合并项目、创建API等操作";
        }
    }

    参考

    https://www.aliyun.com/product/apigateway

    https://www.alibabacloud.com/help/zh/mse/user-guide/configure-http-to-dubbo-protocol-conversion

  • 计算最长路径

    计算最长路径

    微服务场景中很容易出现A->B->C->D->E的情况,现在想找到最长的调用链路,目前已经有从Trace中抽样的数据,遗憾的是只有直接调用关系,比如A->B、C->D、B->C、B->E 这种,所以需要自己加工一下。

    由于场景边界还挺多的,而且不排除后期需要分析其他场景,比如A调用的服务数量等,所以还是考虑使用成熟的工具。

    图形数据库

    第一个想到的就是图形数据库,图形数据库使用图结构进行语义查询的数据库,它使用节点、边和属性来表示和存储数据。这里我们将每个调用拆分为服务作为节点,调用关系作为边来处理。

    为了简单我是用的是Neo4j的嵌入模式,由于一次性需求,每次都新建数据然后查询

    File dataFile = Files.createTempDirectory("neo4j").toFile();
    
    DatabaseManagementService managementService = new DatabaseManagementServiceBuilder(dataFile.toPath()).build();
    GraphDatabaseService graphDb = managementService.database(DEFAULT_DATABASE_NAME);

    为了加快速度,使用ID作为索引

    
    try (Transaction tx = graphDb.beginTx()) {
          Schema schema = tx.schema();
          schema.indexFor(label).on("id").withName("id").create();
          tx.commit();
    }

    然后按需解析数据,添加关联

    Node main = createNode(tx, label, node);
    Node children = createNode(tx, label, callingNode);
    main.createRelationshipTo(children, RelTypes.CALLING);

    然后查询,理论上可以使用Cypher查询语句,由于需求简单我们直接拉出所有路径取最长的

    Traverser traverse = tx.traversalDescription()
                            .relationships(RelTypes.CALLING, Direction.OUTGOING)
                            .evaluator(Evaluators.all()).traverse(mainNode);
    
    for (Path path : traverse) {
         if (longestPath == null || longestPath.length() < path.length()) {
            longestPath = path;
         }
    }

    JGraphT

    Neo4j支持的查询比较多,但是只取最长路径有点大材小用了,而且相对来说性能比较低。JGraphT是一个免费的Java类库,提供数学图形理论对象和算法,也可以用于我们这个场景。

    经过综合衡量,这里使用DirectedPseudograph

    Graph<String, DefaultEdge> graph = new DirectedPseudograph<>(DefaultEdge.class);
    Graphs.addEdgeWithVertices(graph, String.valueOf(node.getLong("id")), String.valueOf(callingNode.getLong("id")));

    JGraphT提供的算法工具大部分是最短路径的,比如AStarShortestPath、BFSShortestPath等,而我们需要最长路径,考虑到图整体不大,我们这里直接用Dijkstra穷举。

    List<List<String>> result = Lists.newArrayList();
    Set<String> allVertices = new HashSet<>(graph.vertexSet());
    allVertices.remove(id);
    
    GraphPath<String, DefaultEdge> longestPath = null;
    double longestPathLength = Double.NEGATIVE_INFINITY;
    
    for (String targetVertex : allVertices) {
         if (targetVertex != null) {
           AllDirectedPaths<String, DefaultEdge> allPaths = new AllDirectedPaths<>(graph);
           List<GraphPath<String, DefaultEdge>> pathsFromNodeA = allPaths.getAllPaths(id, targetVertex, true, null);
    
          for (GraphPath<String, DefaultEdge> path : pathsFromNodeA) {
             if (path.getLength() > longestPathLength) {
                 longestPath = path;
                 longestPathLength = path.getLength();
             }
         }
      }
    }

    当然这个使用实际上有一些限制,比如不能有loop等。

  • 手动运行OpenRewrite配方

    OpenRewrite提供了Maven插件,可以方便运行在Maven管理的项目上,如果需要在其他环境运行,比如自定义的CLI等,就需要手动运行了。

    概念

    先明确下OpenRewrite相关的几个概念

    配方:需要执行的变更,配方可以是内置的,也可以是第三方社区的,更进一步是自定义的。配方的指定主要通过全名完成,比如org.openrewrite.java.OrderImports

    目标项目:要执行变更的项目,一般通过路径表示,也可以通过是远程文件,也可以通过扩展SourceFile适配更多情况

    运行环境:配方执行的基础,提供配方运行、消息处理等,可以扩展ExecutionContext,也可以使用InMemoryExecutionContext

    环境:用于支持运行环境,提供配方管理、资源加载器等,关键类:Environment

    运行流程

    • 一般按照如下流程进行:
    • 初始化环境
    • 激活配方并验证配方
    • 初始化运行环境
    • 解析目标项目,获取源文件集合
    • 运行配方获取变更
    • 根据变更修改项目

    关键代码

    //运行环境准备
    Environment env = Environment.builder().scanRuntimeClasspath().scanUserHome().build();
    ExecutionContext executionContext = initExecutionContext();
    MavenParser.Builder mavenParserBuilder = initMavenRelatedConfig(executionContext);
    
    //指定配方
    Recipe recipe = env.activateRecipes("org.openrewrite.java.RandomizeId");
    Collection<Validated<Object>> validateds = recipe.validateAll();
    //验证配方
     for (Validated<Object> validated : validateds) {
          if (validated instanceof Validated.Invalid) {
            logger.error("recipe validate failed: {}", validated);
             return;
         }
    }
    LargeSourceSet largeSourceSet = new InMemoryLargeSourceSet(sourceFiles);
    
    
    //运行配方
    RecipeRun run = recipe.run(largeSourceSet, executionContext);
    logger.info("Run:{}", run);
    
    for (Result result : run.getChangeset().getAllResults()) {
         logger.info("Result:{}", result);
    }
    
    //写入变更,目前只处理了变更
    for (Result result : run.getChangeset().getAllResults()) {
         if (result.getBefore() != null && result.getAfter() != null) {
             try (final BufferedWriter sourceFileWriter = Files.newBufferedWriter(result.getAfter().getSourcePath())) {
                        sourceFileWriter.write(result.getAfter().printAll());
             }
         }
    }

  • 自定义WildReceipt Paddle Dataset

    自定义WildReceipt Paddle Dataset

    Paddle Dataset是Paddle生态中的数据源抽象,至少需要提供两个方法

    def __getitem__(self, idx):
    
    
    )
    def __len__(self):
    
    )

    将自己的数据封装为Dataset后可以配合高层API使用,也可以享受Paddle生态的各种加强,比如批量加载等。

    Paddle内部包含一些常用的数据集,比如MNIST。常用数据集的封装还包含了自动下载,非常适合新手使用。

    WildReceipt数据集作为文本关键信息提取的基准,无论从数据量还是结构上,都要优于其他公开的数据集。主要用于文档的关键信息提取训练。这里演示下怎么制作Dataset。

    数据集结构

    制作数据集的第一步是了解数据集,了解数据集的结构和需要的输出,这里的输出可能需要关联具体的模型。

    这里使用WildReceipt + SDMGR进行演示。
    WildReceipt数据集主要分两部分,一部分是图片本身,一部分是区域标注。这部分信息存储在txt文件中,图片放在images目录中。由于是在Paddle中,这里直接使用https://paddleocr.bj.bcebos.com/ppstructure/dataset/wildreceipt.tar。

    下面是数据的一些片段

    image_files/Image_12/10/845be0dd6f5b04866a2042abd28d558032ef2576.jpeg	[{"label": "Store_name_value", "transcription": "CHOEUN", "points": [[114.0, 19.0], [230.0, 19.0], [230.0, 1.0], [114.0, 1.0]]}, {"label": "Store_name_value", "transcription": "KOREANRESTAURANT", "points": [[97.0, 35.0], [236.0, 35.0], [236.0, 19.0], [97.0, 19.0]]}, {"label": "Store_addr_value", "transcription": "2621ORANGETHORPEAVE,FULLERTON.", "points": [[29.0, 56.0], [295.0, 56.0], [295.0, 34.0], [29.0, 34.0]]}, {"label": "Tel_value", "transcription": "(714)879-3574", "points": [[48.0, 73.0], [280.0, 73.0], [280.0, 54.0], [48.0, 54.0]]}, {"label": "Others", "transcription": "THANKYOU!!", "points": [[79.0, 92.0], [259.0, 92.0], [259.0, 74.0], [79.0, 74.0]]}, {"label": "Date_key", "transcription": "DATE", "points": [[22.0, 130.0], [61.0, 130.0], [61.0, 112.0], [22.0, 112.0]]}, {"label": "Date_value", "transcription": "12/30/2016FRI", "points": [[70.0, 131.0], [192.0, 131.0], [192.0, 112.0], [70.0, 112.0]]}, {"label": "Time_value", "transcription": "19:19", "points": [[263.0, 128.0], [307.0, 128.0], [307.0, 111.0], [263.0, 111.0]]}, {"label": "Prod_item_value", "transcription": "BIBIM.OCTOPUT1", "points": [[19.0, 168.0], [157.0, 168.0], [157.0, 149.0], [19.0, 149.0]]}, {"label": "Prod_item_value", "transcription": "S-FOODP.CAKT1", "points": [[17.0, 190.0], [158.0, 190.0], [158.0, 171.0], [17.0, 171.0]]}, {"label": "Prod_item_value", "transcription": "PORKDUMPLINT1", "points": [[14.0, 214.0], [158.0, 214.0], [158.0, 192.0], [14.0, 192.0]]}, {"label": "Prod_item_value", "transcription": "LABEEFRIBT1", "points": [[14.0, 236.0], [151.0, 236.0], [151.0, 215.0], [14.0, 215.0]]}, {"label": "Prod_price_value", "transcription": "$13.99", "points": [[254.0, 168.0], [312.0, 168.0], [312.0, 149.0], [254.0, 149.0]]}, {"label": "Prod_price_value", "transcription": "$14.99", "points": [[257.0, 189.0], [314.0, 189.0], [314.0, 170.0], [257.0, 170.0]]}, {"label": "Prod_price_value", "transcription": "$8.99", "points": [[268.0, 212.0], [316.0, 212.0], [316.0, 191.0], [268.0, 191.0]]}, {"label": "Prod_price_value", "transcription": "¥17.99", "points": [[261.0, 234.0], [318.0, 234.0], [318.0, 213.0], [261.0, 213.0]]}, {"label": "Prod_item_key", "transcription": "4.00xITEMS", "points": [[118.0, 260.0], [217.0, 260.0], [217.0, 239.0], [118.0, 239.0]]}, {"label": "Subtotal_key", "transcription": "SUBTOTAL", "points": [[8.0, 285.0], [91.0, 285.0], [91.0, 264.0], [8.0, 264.0]]}, {"label": "Tax_key", "transcription": "TAX1", "points": [[8.0, 312.0], [49.0, 312.0], [49.0, 291.0], [8.0, 291.0]]}, {"label": "Total_key", "transcription": "TOTAL", "points": [[8.0, 336.0], [61.0, 336.0], [61.0, 316.0], [8.0, 316.0]]}, {"label": "Subtotal_value", "transcription": "$55.96", "points": [[263.0, 283.0], [325.0, 283.0], [325.0, 260.0], [263.0, 260.0]]}, {"label": "Tax_value", "transcription": "$4.48", "points": [[274.0, 308.0], [326.0, 308.0], [326.0, 286.0], [274.0, 286.0]]}, {"label": "Total_value", "transcription": "$60.44", "points": [[267.0, 334.0], [328.0, 334.0], [328.0, 310.0], [267.0, 310.0]]}, {"label": "Ignore", "transcription": "", "points": [[269.0, 347.0], [328.0, 347.0], [328.0, 336.0], [269.0, 336.0]]}, {"label": "Ignore", "transcription": "", "points": [[11.0, 347.0], [50.0, 347.0], [50.0, 342.0], [11.0, 342.0]]}, {"label": "Time_key", "transcription": "TIME", "points": [[215.0, 128.0], [253.0, 128.0], [253.0, 112.0], [215.0, 112.0]]}]
    image_files/Image_83/7/f6b397503d69287709ba3872c7e548d45917cd2e.jpeg	[{"label": "Store_name_value", "transcription": "ILIO'S", "points": [[372.0, 242.0], [479.0, 242.0], [479.0, 178.0], [372.0, 178.0]]}, {"label": "Store_name_value", "transcription": "Restaurant", "points": [[338.0, 282.0], [508.0, 282.0], [508.0, 247.0], [338.0, 247.0]]}, {"label": "Store_addr_value", "transcription": "BretonischerRing7", "points": [[285.0, 324.0], [611.0, 324.0], [611.0, 289.0], [285.0, 289.0]]}, {"label": "Store_addr_value", "transcription": "85630Grasbrunn", "points": [[319.0, 367.0], [581.0, 367.0], [581.0, 332.0], [319.0, 332.0]]}, {"label": "Tel_key", "transcription": "TEL:", "points": [[304.0, 409.0], [368.0, 409.0], [368.0, 374.0], [304.0, 374.0]]}, {"label": "Others", "transcription": "Steuer-Nr.:514/78510", "points": [[65.0, 499.0], [442.0, 499.0], [442.0, 462.0], [65.0, 462.0]]}, {"label": "Others", "transcription": "RechnungNr.2844", "points": [[64.0, 623.0], [372.0, 623.0], [372.0, 552.0], [64.0, 552.0]]}, {"label": "Date_key", "transcription": "Datum:", "points": [[64.0, 656.0], [171.0, 656.0], [171.0, 624.0], [64.0, 624.0]]}, {"label": "Date_value", "transcription": "10.07.13", "points": [[197.0, 654.0], [335.0, 654.0], [335.0, 623.0], [197.0, 623.0]]}, {"label": "Time_value", "transcription": "21:52", "points": [[353.0, 653.0], [442.0, 653.0], [442.0, 623.0], [353.0, 623.0]]}, {"label": "Others", "transcription": "Tisch:102/--", "points": [[549.0, 652.0], [792.0, 652.0], [792.0, 617.0], [549.0, 617.0]]}, {"label": "Prod_quantity_value", "transcription": "4x", "points": [[120.0, 742.0], [155.0, 742.0], [155.0, 713.0], [120.0, 713.0]]}, {"label": "Prod_quantity_value", "transcription": "8x", "points": [[118.0, 788.0], [154.0, 788.0], [154.0, 756.0], [118.0, 756.0]]}, {"label": "Prod_quantity_value", "transcription": "3x", "points": [[118.0, 831.0], [151.0, 831.0], [151.0, 801.0], [118.0, 801.0]]}, {"label": "Prod_quantity_value", "transcription": "1x", "points": [[119.0, 876.0], [152.0, 876.0], [152.0, 845.0], [119.0, 845.0]]}, {"label": "Prod_quantity_value", "transcription": "1x", "points": [[119.0, 923.0], [152.0, 923.0], [152.0, 890.0], [119.0, 890.0]]}, {"label": "Prod_quantity_value", "transcription": "1x", "points": [[119.0, 967.0], [152.0, 967.0], [152.0, 936.0], [119.0, 936.0]]}, {"label": "Prod_quantity_value", "transcription": "1x", "points": [[118.0, 1012.0], [151.0, 1012.0], [151.0, 981.0], [118.0, 981.0]]}, {"label": "Prod_quantity_value", "transcription": "1x", "points": [[118.0, 1058.0], [149.0, 1058.0], [149.0, 1027.0], [118.0, 1027.0]]}, {"label": "Prod_quantity_value", "transcription": "2x", "points": [[112.0, 1104.0], [147.0, 1104.0], [147.0, 1070.0], [112.0, 1070.0]]}, {"label": "Prod_quantity_value", "transcription": "1x", "points": [[115.0, 1151.0], [145.0, 1151.0], [145.0, 1119.0], [115.0, 1119.0]]}, {"label": "Prod_quantity_value", "transcription": "1x", "points": [[115.0, 1198.0], [146.0, 1198.0], [146.0, 1164.0], [115.0, 1164.0]]}, {"label": "Prod_quantity_value", "transcription": "1x", "points": [[114.0, 1243.0], [147.0, 1243.0], [147.0, 1210.0], [114.0, 1210.0]]}, {"label": "Prod_quantity_value", "transcription": "1x", "points": [[113.0, 1288.0], [147.0, 1288.0], [147.0, 1253.0], [113.0, 1253.0]]}, {"label": "Prod_quantity_value", "transcription": "1x", "points": [[113.0, 1331.0], [145.0, 1331.0], [145.0, 1299.0], [113.0, 1299.0]]}, {"label": "Prod_item_value", "transcription": "Tee", "points": [[165.0, 1332.0], [218.0, 1332.0], [218.0, 1298.0], [165.0, 1298.0]]}, {"label": "Prod_item_value", "transcription": "Stifado", "points": [[165.0, 1285.0], [292.0, 1285.0], [292.0, 1251.0], [165.0, 1251.0]]}, {"label": "Prod_item_value", "transcription": "SchweinefiletMeta", "points": [[165.0, 1238.0], [493.0, 1238.0], [493.0, 1204.0], [165.0, 1204.0]]}, {"label": "Prod_item_value", "transcription": "BiftekiMetaxa", "points": [[165.0, 1193.0], [419.0, 1193.0], [419.0, 1159.0], [165.0, 1159.0]]}, {"label": "Ignore", "transcription": "", "points": [[165.0, 1152.0], [440.0, 1152.0], [440.0, 1111.0], [165.0, 1111.0]]}, {"label": "Prod_item_value", "transcription": "GyrosFolie", "points": [[167.0, 1107.0], [370.0, 1107.0], [370.0, 1066.0], [167.0, 1066.0]]}, {"label": "Prod_item_value", "transcription": "BabyKalamariGefu", "points": [[167.0, 1062.0], [495.0, 1062.0], [495.0, 1019.0], [167.0, 1019.0]]}, {"label": "Prod_item_value", "transcription": "Gyros", "points": [[168.0, 1016.0], [260.0, 1016.0], [260.0, 979.0], [168.0, 979.0]]}, {"label": "Prod_item_value", "transcription": "VegetarischeVaria", "points": [[171.0, 971.0], [493.0, 971.0], [493.0, 931.0], [171.0, 931.0]]}, {"label": "Prod_item_value", "transcription": "GrossesWasser", "points": [[169.0, 922.0], [422.0, 922.0], [422.0, 889.0], [169.0, 889.0]]}, {"label": "Prod_item_value", "transcription": "Saft0,25", "points": [[171.0, 877.0], [336.0, 877.0], [336.0, 841.0], [171.0, 841.0]]}, {"label": "Prod_item_value", "transcription": "Hefe-Weissbier", "points": [[171.0, 833.0], [422.0, 833.0], [422.0, 795.0], [171.0, 795.0]]}, {"label": "Prod_item_value", "transcription": "Weissbierdunkel", "points": [[172.0, 788.0], [455.0, 788.0], [455.0, 750.0], [172.0, 750.0]]}, {"label": "Prod_item_value", "transcription": "LowenbrauOriginal", "points": [[173.0, 742.0], [490.0, 742.0], [490.0, 708.0], [173.0, 708.0]]}, {"label": "Others", "transcription": "a", "points": [[511.0, 738.0], [527.0, 738.0], [527.0, 713.0], [511.0, 713.0]]}, {"label": "Others", "transcription": "a", "points": [[512.0, 782.0], [527.0, 782.0], [527.0, 758.0], [512.0, 758.0]]}, {"label": "Others", "transcription": "a", "points": [[512.0, 826.0], [529.0, 826.0], [529.0, 804.0], [512.0, 804.0]]}, {"label": "Others", "transcription": "a", "points": [[511.0, 1098.0], [527.0, 1098.0], [527.0, 1073.0], [511.0, 1073.0]]}, {"label": "Others", "transcription": "9,90", "points": [[564.0, 1101.0], [635.0, 1101.0], [635.0, 1066.0], [564.0, 1066.0]]}, {"label": "Others", "transcription": "3,30", "points": [[564.0, 829.0], [632.0, 829.0], [632.0, 795.0], [564.0, 795.0]]}, {"label": "Others", "transcription": "3,30", "points": [[564.0, 785.0], [633.0, 785.0], [633.0, 751.0], [564.0, 751.0]]}, {"label": "Others", "transcription": "3,00", "points": [[566.0, 743.0], [635.0, 743.0], [635.0, 707.0], [566.0, 707.0]]}, {"label": "Prod_price_value", "transcription": "12,00", "points": [[691.0, 742.0], [776.0, 742.0], [776.0, 706.0], [691.0, 706.0]]}, {"label": "Prod_price_value", "transcription": "26,40", "points": [[687.0, 786.0], [776.0, 786.0], [776.0, 750.0], [687.0, 750.0]]}, {"label": "Prod_price_value", "transcription": "9,90", "points": [[706.0, 830.0], [778.0, 830.0], [778.0, 795.0], [706.0, 795.0]]}, {"label": "Prod_price_value", "transcription": "2,50", "points": [[706.0, 873.0], [779.0, 873.0], [779.0, 840.0], [706.0, 840.0]]}, {"label": "Prod_price_value", "transcription": "2,40", "points": [[706.0, 922.0], [780.0, 922.0], [780.0, 885.0], [706.0, 885.0]]}, {"label": "Prod_price_value", "transcription": "9,90", "points": [[707.0, 967.0], [778.0, 967.0], [778.0, 931.0], [707.0, 931.0]]}, {"label": "Prod_price_value", "transcription": "8,90", "points": [[706.0, 1014.0], [780.0, 1014.0], [780.0, 976.0], [706.0, 976.0]]}, {"label": "Prod_price_value", "transcription": "12,90", "points": [[693.0, 1059.0], [780.0, 1059.0], [780.0, 1022.0], [693.0, 1022.0]]}, {"label": "Prod_price_value", "transcription": "19,80", "points": [[694.0, 1105.0], [781.0, 1105.0], [781.0, 1069.0], [694.0, 1069.0]]}, {"label": "Prod_price_value", "transcription": "6,90", "points": [[708.0, 1150.0], [782.0, 1150.0], [782.0, 1114.0], [708.0, 1114.0]]}, {"label": "Prod_price_value", "transcription": "11,90", "points": [[696.0, 1196.0], [783.0, 1196.0], [783.0, 1160.0], [696.0, 1160.0]]}, {"label": "Prod_price_value", "transcription": "13,90", "points": [[697.0, 1242.0], [784.0, 1242.0], [784.0, 1206.0], [697.0, 1206.0]]}, {"label": "Prod_price_value", "transcription": "14,90", "points": [[696.0, 1289.0], [785.0, 1289.0], [785.0, 1253.0], [696.0, 1253.0]]}, {"label": "Prod_price_value", "transcription": "2,10", "points": [[711.0, 1336.0], [784.0, 1336.0], [784.0, 1299.0], [711.0, 1299.0]]}, {"label": "Others", "transcription": "1", "points": [[807.0, 1333.0], [818.0, 1333.0], [818.0, 1301.0], [807.0, 1301.0]]}, {"label": "Others", "transcription": "1", "points": [[807.0, 1287.0], [818.0, 1287.0], [818.0, 1254.0], [807.0, 1254.0]]}, {"label": "Others", "transcription": "1", "points": [[805.0, 1241.0], [817.0, 1241.0], [817.0, 1210.0], [805.0, 1210.0]]}, {"label": "Others", "transcription": "1", "points": [[804.0, 1195.0], [816.0, 1195.0], [816.0, 1163.0], [804.0, 1163.0]]}, {"label": "Others", "transcription": "1", "points": [[804.0, 1150.0], [816.0, 1150.0], [816.0, 1118.0], [804.0, 1118.0]]}, {"label": "Others", "transcription": "1", "points": [[805.0, 1104.0], [814.0, 1104.0], [814.0, 1073.0], [805.0, 1073.0]]}, {"label": "Others", "transcription": "1", "points": [[804.0, 1055.0], [814.0, 1055.0], [814.0, 1024.0], [804.0, 1024.0]]}, {"label": "Others", "transcription": "1", "points": [[802.0, 1011.0], [814.0, 1011.0], [814.0, 979.0], [802.0, 979.0]]}, {"label": "Others", "transcription": "1", "points": [[801.0, 964.0], [813.0, 964.0], [813.0, 932.0], [801.0, 932.0]]}, {"label": "Others", "transcription": "1", "points": [[800.0, 918.0], [812.0, 918.0], [812.0, 887.0], [800.0, 887.0]]}, {"label": "Others", "transcription": "1", "points": [[801.0, 872.0], [812.0, 872.0], [812.0, 840.0], [801.0, 840.0]]}, {"label": "Others", "transcription": "1", "points": [[801.0, 827.0], [811.0, 827.0], [811.0, 795.0], [801.0, 795.0]]}, {"label": "Others", "transcription": "1", "points": [[799.0, 781.0], [811.0, 781.0], [811.0, 749.0], [799.0, 749.0]]}, {"label": "Others", "transcription": "1", "points": [[798.0, 737.0], [809.0, 737.0], [809.0, 705.0], [798.0, 705.0]]}, {"label": "Subtotal_key", "transcription": "Netto(1)", "points": [[53.0, 1428.0], [200.0, 1428.0], [200.0, 1388.0], [53.0, 1388.0]]}, {"label": "Tax_key", "transcription": "+19,0%MwSt:", "points": [[52.0, 1474.0], [303.0, 1474.0], [303.0, 1436.0], [52.0, 1436.0]]}, {"label": "Subtotal_value", "transcription": "Eur129,75", "points": [[252.0, 1427.0], [473.0, 1427.0], [473.0, 1390.0], [252.0, 1390.0]]}, {"label": "Tax_value", "transcription": "24,65", "points": [[380.0, 1478.0], [471.0, 1478.0], [471.0, 1438.0], [380.0, 1438.0]]}, {"label": "Total_key", "transcription": "Summe:", "points": [[40.0, 1603.0], [148.0, 1603.0], [148.0, 1535.0], [40.0, 1535.0]]}, {"label": "Others", "transcription": "EsbedienteSie:George", "points": [[37.0, 1654.0], [469.0, 1654.0], [469.0, 1611.0], [37.0, 1611.0]]}, {"label": "Total_value", "transcription": "Eur154,40", "points": [[565.0, 1617.0], [788.0, 1617.0], [788.0, 1537.0], [565.0, 1537.0]]}, {"label": "Tel_value", "transcription": "089-46169340", "points": [[386.0, 409.0], [599.0, 409.0], [599.0, 373.0], [386.0, 373.0]]}]

    每一行是一条数据,数据前部分是图片路径,后半部分是标注信息,是一段JSON。JSON主要是三部分数据,一是标签,二是文本内容,三是区域。

    这里的label还需要配合一个class_list使用,另外根据文字内容还需要对应的字典。

    /
    \
    .
    $
    £
    €
    ¥
    :
    -
    ,
    *
    #
    (
    )
    %
    @
    !
    '
    &
    =
    >
    +
    "
    Ignore
    Store_name_value
    Store_name_key
    Store_addr_value
    Store_addr_key
    Tel_value
    Tel_key
    Date_value
    Date_key
    Time_value
    Time_key
    Prod_item_value
    Prod_item_key
    Prod_quantity_value
    Prod_quantity_key
    Prod_price_value
    Prod_price_key
    Subtotal_value
    Subtotal_key
    Tax_value
    Tax_key
    Tips_value
    Tips_key
    Total_value
    Total_key
    Others

    下载数据

    Paddle内置了一个下载方法,直接使用下载文件

    class WildReceiptDataset(paddle.io.Dataset):
        NAME = "wildreceipt"
        DATASET_URL = "https://paddleocr.bj.bcebos.com/ppstructure/dataset/wildreceipt.tar"
        DATASET_MD5 = "0b9abbc025e85515247f8a464c7b44dc"
    
        def __init__(self, path=None, mode="train", transform=None, download=True):
            super(WildReceiptDataset, self).__init__()
    
            assert mode.lower() in [
                "train",
                "test",
            ], f"mode should be 'train' or 'test', but got {mode}"
    
            self.mode = mode.lower()
            self.path = path
            if self.path is None:
                assert (
                    download
                ), "image_path is not set and downloading automatically is disabled"
                self.path = paddle.dataset.common.download(
                    self.DATASET_URL, self.NAME, self.DATASET_MD5
                )
    
            self.transform = transform

    数据会自动缓存,避免重复下载

    解析数据

    我们先简单获取图片和label信息,label信息直接返回json

    class WildReceiptDataset(paddle.io.Dataset):
        def _parse_dataset(self, buffer_size=100):
            self.images = []
            self.labels = []
    
            main_file = "wildreceipt/wildreceipt_" + self.mode + ".txt"
    
            with tarfile.open(self.path) as tarFile:
                member = tarFile.getmember(main_file)
                f = tarFile.extractfile(member)
                if f is not None:
                    content = f.read()
    
                text = content.decode("utf-8")
                lines = text.split("\n")
    
                for line in lines:
                    if line == "":
                        continue
                    substr = line.split("\t")
                    file_name = substr[0]
                    label = substr[1]
    
                    self.images.append(file_name)
                    self.labels.append(label)

    然后在获取数据的时候解析图片

    class WildReceiptDataset(paddle.io.Dataset):
    NAME = “wildreceipt”
    DATASET_URL = “https://paddleocr.bj.bcebos.com/ppstructure/dataset/wildreceipt.tar”
    DATASET_MD5 = “0b9abbc025e85515247f8a464c7b44dc”

    class WildReceiptDataset(paddle.io.Dataset):
    
    
        def __getitem__(self, idx):
            file_name = self.images[idx]
            with tarfile.open(self.path) as tarFile:
                image_data = io.BytesIO(
                    tarFile.extractfile("wildreceipt/" + file_name).read()
                )
                image = Image.open(image_data)
            return image, self.labels[idx]