最新国产好看的视频,伊人天堂AV在线,国产Aaaaaa视频,蜜臀视频在线观看一区,人妻av色图,密臀久久久精品影片,青青视频免费观看毛片,久草在线观看视,国产三级精品色情在线

pytorch單元測(cè)試的實(shí)現(xiàn)示例

 更新時(shí)間:2024年04月18日 10:57:40   作者:Hi20240217  
單元測(cè)試是一種軟件測(cè)試方法,本文主要介紹了pytorch單元測(cè)試的實(shí)現(xiàn)示例,文中通過(guò)示例代碼介紹的非常詳細(xì),對(duì)大家的學(xué)習(xí)或者工作具有一定的參考學(xué)習(xí)價(jià)值,需要的朋友們下面隨著小編來(lái)一起學(xué)習(xí)學(xué)習(xí)吧

希望測(cè)試pytorch各種算子、block、網(wǎng)絡(luò)等在不同硬件平臺(tái),不同軟件版本下的計(jì)算誤差、耗時(shí)、內(nèi)存占用等指標(biāo).

本文基于torch.testing._internal

一.公共模塊[common.py]

import torch
from torch import nn
import math
import torch.nn.functional as F
import time
import os
import socket
import sys
from datetime import datetime
import numpy as np
import collections
import math
import json
import copy
import traceback
import subprocess
import unittest
import torch
import inspect
from torch.testing._internal.common_utils import TestCase, run_tests,parametrize,instantiate_parametrized_tests
from torch.testing._internal.common_distributed import MultiProcessTestCase
import torch.distributed as dist

os.environ["MASTER_ADDR"] = "localhost"
os.environ["MASTER_PORT"] = "29500"
os.environ["RANDOM_SEED"] = "0" 

device="cpu"
device_type="cpu"
device_name="cpu"

try:
    if torch.cuda.is_available():     
        device_name=torch.cuda.get_device_name().replace(" ","")
        device="cuda:0"
        device_type="cuda"
        ccl_backend='nccl'
except:
    pass

host_name=socket.gethostname()    
sdk_version=os.getenv("SDK_VERSION","")   						 #從環(huán)境變量中獲取sdk版本號(hào)
metric_data_root=os.getenv("TORCH_UT_METRICS_DATA","./ut_data")  #日志存放的目錄
device_count=torch.cuda.device_count()

if not os.path.exists(metric_data_root):
    os.makedirs(metric_data_root)

def device_warmup(device):
    '''設(shè)備warmup,確保設(shè)備已經(jīng)正常工作,排除設(shè)備初始化的耗時(shí)'''
    left = torch.rand([128,512], dtype = torch.float16).to(device)
    right = torch.rand([512,128], dtype = torch.float16).to(device)
    out=torch.matmul(left,right)
    torch.cuda.synchronize()

torch.manual_seed(1) 
np.random.seed(1)

def loop_decorator(loops,rank=0):
    '''循環(huán)裝飾器,用于統(tǒng)計(jì)函數(shù)的執(zhí)行時(shí)間,內(nèi)存占用等'''
    def decorator(func):
        def wrapper(*args,**kwargs):
            latency=[]
            memory_allocated_t0=torch.cuda.memory_allocated(rank)
            for _ in range(loops):
                input_copy=[x.clone() for x in args]
                beg= datetime.now().timestamp() * 1e6
                pred= func(*input_copy)
                gt=kwargs["golden"]
                torch.cuda.synchronize()
                end=datetime.now().timestamp() * 1e6
                mse = torch.mean(torch.pow(pred.cpu().float()- gt.cpu().float(), 2)).item()
                latency.append(end-beg)
            memory_allocated_t1=torch.cuda.memory_allocated(rank)
            avg_latency=np.mean(latency[len(latency)//2:]).round(3)
            first_latency=latency[0]
            return { "first_latency":first_latency,"avg_latency":avg_latency,
                      "memory_allocated":memory_allocated_t1-memory_allocated_t0,
                      "mse":mse}
        return wrapper
    return decorator

class TorchUtMetrics:
    '''用于統(tǒng)計(jì)測(cè)試結(jié)果,比較之前的最小值'''
    def __init__(self,ut_name,thresold=0.2,rank=0):
        self.ut_name=f"{ut_name}_{rank}"
        self.thresold=thresold
        self.rank=rank
        self.data={"ut_name":self.ut_name,"metrics":[]}
        self.metrics_path=os.path.join(metric_data_root,f"{self.ut_name}_{self.rank}.jon")
        try:
            with open(self.metrics_path,"r") as f:
                self.data=json.loads(f.read())
        except:
            pass

    def __enter__(self):
        self.beg= datetime.now().timestamp() * 1e6
        return self

    def __exit__(self, exc_type, exc_val, exc_tb):        
        self.report()
        self.save_data()

    def save_data(self):
        with open(self.metrics_path,"w") as f:
            f.write(json.dumps(self.data,indent=4))

    def set_metrics(self,metrics):
        self.end=datetime.now().timestamp() * 1e6
        item=collections.OrderedDict()
        item["time"]=datetime.now().strftime('%Y-%m-%d %H:%M:%S.%f')
        item["sdk_version"]=sdk_version
        item["device_name"]=device_name
        item["host_name"]=host_name
        item["metrics"]=metrics
        item["metrics"]["e2e_time"]=self.end-self.beg
        self.cur_item=item
        self.data["metrics"].append(self.cur_item)

    def get_metric_names(self):
        return self.data["metrics"][0]["metrics"].keys()

    def get_min_metric(self,metric_name,devicename=None):
        min_value=0
        min_value_index=-1
        for idx,item in enumerate(self.data["metrics"]):
            if devicename and (devicename!=item['device_name']):                
                continue            
            val=float(item["metrics"][metric_name])
            if min_value_index==-1 or val<min_value:
                min_value=val
                min_value_index=idx
        return min_value,min_value_index

    def get_metric_info(self,index):
        metrics=self.data["metrics"][index]
        return f'{metrics["device_name"]}@{metrics["sdk_version"]}'

    def report(self):
        assert len(self.data["metrics"])>0
        for metric_name in self.get_metric_names():
            min_value,min_value_index=self.get_min_metric(metric_name)
            min_value_same_dev,min_value_index_same_dev=self.get_min_metric(metric_name,device_name)
            cur_value=float(self.cur_item["metrics"][metric_name])
            print(f"-------------------------------{metric_name}-------------------------------")
            print(f"{cur_value}#{device_name}@{sdk_version}")
            if min_value_index_same_dev>=0:
                print(f"{min_value_same_dev}#{self.get_metric_info(min_value_index_same_dev)}")
            if min_value_index>=0:
                print(f"{min_value}#{self.get_metric_info(min_value_index)}")

二.普通算子測(cè)試[test_clone.py]

from common import *
class TestCaseClone(TestCase):
    #如果不滿(mǎn)足條件,則跳過(guò)這個(gè)測(cè)試
    @unittest.skipIf(device_count>1, "Not enough devices") 
    def test_todo(self):
        print(".TODO")

    #框架會(huì)自動(dòng)遍歷以下參數(shù)組合
    @parametrize("shape", [(10240,20480),(128,256)])
    @parametrize("dtype", [torch.float16,torch.float32])
    def test_clone(self,shape,dtype):
        
        #讓這個(gè)函數(shù)循環(huán)執(zhí)行l(wèi)oops次,統(tǒng)計(jì)第一次執(zhí)行的耗時(shí)、后半段的平均時(shí)間、整個(gè)執(zhí)行過(guò)程總的GPU內(nèi)存使用量
        @loop_decorator(loops=5)
        def run(input_dev):
            output=input_dev.clone()
            return output
        
        #記錄整個(gè)測(cè)試的總耗時(shí),保存統(tǒng)計(jì)量,輸出摘要(self._testMethodName:測(cè)試方法,result:函數(shù)返回值,metrics:統(tǒng)計(jì)量)
        with TorchUtMetrics(ut_name=self._testMethodName,thresold=0.2) as m:
            input_host=torch.ones(shape,dtype=dtype)*np.random.rand()
            input_dev=input_host.to(device)
            metrics=run(input_dev,golden=input_host.cpu())
            m.set_metrics(metrics)
            assert(metrics["mse"]==0)
        
instantiate_parametrized_tests(TestCaseClone)

if __name__ == "__main__":
    run_tests()

三.集合通信測(cè)試[test_ccl.py]

from common import *
class TestCCL(MultiProcessTestCase):
    '''CCL測(cè)試用例'''
    def _create_process_group_vccl(self, world_size, store):
        dist.init_process_group(
            ccl_backend, world_size=world_size, rank=self.rank, store=store
        )        
        pg = dist.distributed_c10d._get_default_group()
        return pg

    def setUp(self):
        super().setUp()
        self._spawn_processes()

    def tearDown(self):
        super().tearDown()
        try:
            os.remove(self.file_name)
        except OSError:
            pass

    @property
    def world_size(self):
        return 4
      
    #框架會(huì)自動(dòng)遍歷以下參數(shù)組合
    @unittest.skipIf(device_count<4, "Not enough devices") 
    @parametrize("op",[dist.ReduceOp.SUM])
    @parametrize("shape", [(1024,8192)])
    @parametrize("dtype", [torch.int64])
    def test_allreduce(self,op,shape,dtype):
        if self.rank >= self.world_size:
            return
        
        store = dist.FileStore(self.file_name, self.world_size)
        pg = self._create_process_group_vccl(self.world_size, store)
        if not torch.distributed.is_initialized():
            return
    
        torch.cuda.set_device(self.rank)
        device = torch.device(device_type,self.rank)
        device_warmup(device)
        #讓這個(gè)函數(shù)循環(huán)執(zhí)行l(wèi)oops次,統(tǒng)計(jì)第一次執(zhí)行的耗時(shí)、后半段的平均時(shí)間、整個(gè)執(zhí)行過(guò)程總的GPU內(nèi)存使用量
        @loop_decorator(loops=5,rank=self.rank)
        def run(input_dev):
            dist.all_reduce(input_dev, op=op)
            return input_dev
        
        #記錄整個(gè)測(cè)試的總耗時(shí),保存統(tǒng)計(jì)量,輸出摘要(self._testMethodName:測(cè)試方法,result:函數(shù)返回值,metrics:統(tǒng)計(jì)量)
        with TorchUtMetrics(ut_name=self._testMethodName,thresold=0.2,rank=self.rank) as m:
            input_host=torch.ones(shape,dtype=dtype)*(100+self.rank)
            gt=[torch.ones(shape,dtype=dtype)*(100+i) for i in range(self.world_size)]
            gt_=gt[0]
            for i in range(1,self.world_size):
                gt_=gt_+gt[i]
            input_dev=input_host.to(device)
            metrics=run(input_dev,golden=gt_)
            m.set_metrics(metrics)
            assert(metrics["mse"]==0)
        dist.destroy_process_group(pg)
    
instantiate_parametrized_tests(TestCCL)

if __name__ == "__main__":
    run_tests()

四.測(cè)試命令

# 運(yùn)行所有的測(cè)試
pytest -v -s -p no:warnings --html=torch_report.html --self-contained-html --capture=sys ./

# 運(yùn)行某一個(gè)測(cè)試
python3 test_clone.py -k "test_clone_shape_(128, 256)_float32"

五.測(cè)試報(bào)告

在這里插入圖片描述

到此這篇關(guān)于pytorch單元測(cè)試的實(shí)現(xiàn)示例的文章就介紹到這了,更多相關(guān)pytorch單元測(cè)試內(nèi)容請(qǐng)搜索腳本之家以前的文章或繼續(xù)瀏覽下面的相關(guān)文章希望大家以后多多支持腳本之家! 

相關(guān)文章

  • 用python進(jìn)行線(xiàn)性/非線(xiàn)性擬合的三種方法

    用python進(jìn)行線(xiàn)性/非線(xiàn)性擬合的三種方法

    這篇文章主要給大家介紹了關(guān)于用python進(jìn)行線(xiàn)性/非線(xiàn)性擬合的三種方法,數(shù)據(jù)分析中經(jīng)常會(huì)使用到數(shù)據(jù)擬合,文中通過(guò)實(shí)例代碼以及圖文介紹的非常詳細(xì),需要的朋友可以參考下
    2023-07-07
  • Python3使用SMTP發(fā)送帶附件郵件

    Python3使用SMTP發(fā)送帶附件郵件

    這篇文章主要為大家詳細(xì)介紹了Python3使用SMTP發(fā)送帶附件郵件,文中示例代碼介紹的非常詳細(xì),具有一定的參考價(jià)值,感興趣的小伙伴們可以參考一下
    2018-06-06
  • Python中type()函數(shù)的具體使用

    Python中type()函數(shù)的具體使用

    在Python中,type()函數(shù)是一個(gè)非常有用的工具,它可以查看變量或?qū)ο蟮臄?shù)據(jù)類(lèi)型,本文主要介紹了Python中type()函數(shù)的具體使用,感興趣的可以一起來(lái)了解一下
    2024-01-01
  • Python實(shí)現(xiàn)DBSCAN聚類(lèi)算法并樣例測(cè)試

    Python實(shí)現(xiàn)DBSCAN聚類(lèi)算法并樣例測(cè)試

    聚類(lèi)是一種機(jī)器學(xué)習(xí)技術(shù),它涉及到數(shù)據(jù)點(diǎn)的分組,聚類(lèi)是一種無(wú)監(jiān)督學(xué)習(xí)的方法,是許多領(lǐng)域中常用的統(tǒng)計(jì)數(shù)據(jù)分析技術(shù)。本文給大家分享Python實(shí)現(xiàn)DBSCAN聚類(lèi)算法并樣例測(cè)試,感興趣的朋友一起看看吧
    2021-06-06
  • 淺談python3 構(gòu)造函數(shù)和析構(gòu)函數(shù)

    淺談python3 構(gòu)造函數(shù)和析構(gòu)函數(shù)

    這篇文章主要介紹了淺談python3 構(gòu)造函數(shù)和析構(gòu)函數(shù),具有很好的參考價(jià)值,希望對(duì)大家有所幫助。一起跟隨小編過(guò)來(lái)看看吧
    2020-03-03
  • Python多進(jìn)程原理與用法分析

    Python多進(jìn)程原理與用法分析

    這篇文章主要介紹了Python多進(jìn)程原理與用法,結(jié)合實(shí)例形式分析了Python多進(jìn)程原理、開(kāi)啟使用進(jìn)程、進(jìn)程隊(duì)列、進(jìn)程池等相關(guān)概念與使用方法,需要的朋友可以參考下
    2018-08-08
  • python字典進(jìn)行運(yùn)算原理及實(shí)例分享

    python字典進(jìn)行運(yùn)算原理及實(shí)例分享

    在本篇文章里小編給大家整理的是一篇關(guān)于python字典進(jìn)行運(yùn)算原理及實(shí)例分享內(nèi)容,有需要的朋友們可以測(cè)試下。
    2021-08-08
  • Python如何使用vars返回對(duì)象的屬性列表

    Python如何使用vars返回對(duì)象的屬性列表

    這篇文章主要介紹了Python如何使用vars返回對(duì)象的屬性列表,文中通過(guò)示例代碼介紹的非常詳細(xì),對(duì)大家的學(xué)習(xí)或者工作具有一定的參考學(xué)習(xí)價(jià)值,需要的朋友可以參考下
    2020-10-10
  • pytorch實(shí)現(xiàn)手寫(xiě)數(shù)字圖片識(shí)別

    pytorch實(shí)現(xiàn)手寫(xiě)數(shù)字圖片識(shí)別

    這篇文章主要為大家詳細(xì)介紹了pytorch實(shí)現(xiàn)手寫(xiě)數(shù)字圖片識(shí)別,文中示例代碼介紹的非常詳細(xì),具有一定的參考價(jià)值,感興趣的小伙伴們可以參考一下
    2021-05-05
  • Python?SQLAlchemy插入日期時(shí)間時(shí)區(qū)詳解

    Python?SQLAlchemy插入日期時(shí)間時(shí)區(qū)詳解

    SQLAlchemy是一個(gè)功能強(qiáng)大且流行的?Python?庫(kù),它提供了一種靈活有效的與數(shù)據(jù)庫(kù)交互的方式,在本文中,我們將了解SQLAlchemy如何更新日期、時(shí)間和時(shí)區(qū)并將其插入數(shù)據(jù)庫(kù),感興趣的可以了解下
    2023-09-09

最新評(píng)論

花莲县| 楚雄市| 颍上县| 清苑县| 吴堡县| 荔浦县| 镇平县| 宜宾市| 桂阳县| 芜湖市| 清原| 杭州市| 平遥县| 武鸣县| 五原县| 水城县| 石嘴山市| 施甸县| 陇南市| 新龙县| 江孜县| 万全县| 虹口区| 大邑县| 永德县| 银川市| 乌拉特中旗| 东山县| 岑溪市| 始兴县| 宁强县| 台东市| 肃宁县| 焉耆| 巴彦县| 肥西县| 无为县| 南汇区| 万荣县| 延寿县| 酒泉市|