【MindStudio训练营第一期】垃圾分类案例代码详解
收藏回复举报
【MindStudio训练营第一期】垃圾分类案例代码详解
发表于2022-12-30 21:02:15
0 查看

基于垃圾分类案例

直接看classify_test.py中的main函数了解流程:

acl_resource = AclLiteResource() #创建一个AclResource类的实例
    acl_resource.init() ##AscendCL资源初始化(申请device、创建Context、创建Stream)
    model = AclLiteModel(MODEL_PATH)

这里主要做的事情:

  1. 运行管理资源申请:用于初始化系统内部资源,固定的调用流程。
  2. 加载模型文件并构建输出内存

之后是读取需要处理的图像

dvpp = AclLiteImageProc(acl_resource)

# 之后使用DVPP处理图像的代码
 resized_image = pre_process(image, dvpp)

# 这个函数我们再看看它做了什么
def pre_process(image, dvpp):
    """preprocess"""
    image_input = image.copy_to_dvpp()
    yuv_image = dvpp.jpegd(image_input)

    print("decode jpeg end")
    resized_image = dvpp.resize(yuv_image, 
                    MODEL_WIDTH, MODEL_HEIGHT)

    print("resize yuv end")
    return resized_image
# 不难看出,做了两件事,使用DVPP的JPEGD接口解码图片,
# JPEGD解码出来的是YUV格式。然后将图片进行缩放

# 其中resize的尺寸由全局变量得到
MODEL_WIDTH = 224
MODEL_HEIGHT = 224

这里使用DVPP的JPEGD接口解码图片,JPEGD解码出来的是YUV格式。然后将图片进行缩放。

模型推理

# 用处理后的图片作为输入,使用模型进行推理并得到推理结果
result = model.execute([resized_image,])

后处理 得到模型推理结果后只是一些向量数字,数字的大小代表了属于不同类的概率,选择最大置信度即可。

post_process(result, image_file)    

def post_process(infer_output, image_file):
    print("post process")
    data = infer_output[0]
    vals = data.flatten()
    top_k = vals.argsort()[-1:-6:-1] # 排序结果
    object_class = get_image_net_class(top_k[0]) #选择置信度最高的结果
    output_path = os.path.join(os.path.join(SRC_PATH, "../out"), os.path.basename(image_file))
    origin_image = Image.open(image_file)
    draw = ImageDraw.Draw(origin_image)
    font = ImageFont.truetype("/usr/share/fonts/truetype/dejavu/DejaVuSans-Bold.ttf", size=20)
    font.size =50
    draw.text((10, 50), object_class, font=font, fill=255) #将预测的结果放置到图片上
    origin_image.save(output_path)
    object_class = get_image_net_class(top_k[0])        
    return

我要发帖子