特性分析)
Format 推導(dǎo)Infer Format特性分析【免費(fèi)下載鏈接】geGEGraph Engine是面向昇騰的圖編譯器和執(zhí)行器提供了計(jì)算圖優(yōu)化、多流并行、內(nèi)存復(fù)用和模型下沉等技術(shù)手段加速模型執(zhí)行效率減少模型內(nèi)存占用。 GE 提供對(duì) PyTorch、TensorFlow 前端的友好接入能力并同時(shí)支持 onnx、pb 等主流模型格式的解析與編譯。項(xiàng)目地址: https://gitcode.com/cann/ge1. 特性背景1.1 問(wèn)題的本質(zhì)深度學(xué)習(xí)框架PyTorch、TensorFlow 等在構(gòu)造計(jì)算圖時(shí)用戶關(guān)注的是計(jì)算語(yǔ)義——張量的維度、算子的數(shù)學(xué)含義以及數(shù)據(jù)依賴(lài)關(guān)系。但昇騰 AI 處理器Ascend NPU的硬件架構(gòu)對(duì)數(shù)據(jù)的內(nèi)存布局有特定要求例如Conv2D 的圖片輸入在硬件上親和 NC1HWC0 格式將 C 軸按 16 對(duì)齊拆分MatMul 的權(quán)重親和 FRACTAL_NZ 格式不同算子對(duì)格式有不同的支持能力和性能偏好用戶以 NCHW 或 NHWC 等通用格式描述的模型在昇騰設(shè)備上執(zhí)行時(shí)需要被轉(zhuǎn)換為硬件親和的內(nèi)存布局。這涉及兩個(gè)核心問(wèn)題語(yǔ)義理解如何正確還原用戶在整張計(jì)算圖中表達(dá)的格式語(yǔ)義執(zhí)行優(yōu)化如何為算子選擇合適的執(zhí)行格式并盡量減少數(shù)據(jù)重排TransData的開(kāi)銷(xiāo)1.2 為什么需要兩套格式字段GE 引入了Origin Format和Storage Format也稱(chēng) Running Format兩套表示體系其根本原因是語(yǔ)義正確性和執(zhí)行效率需要獨(dú)立建模視角字段職責(zé)來(lái)源Originorigin_format_表達(dá)用戶構(gòu)造計(jì)算圖時(shí)的原始格式語(yǔ)義前端框架或用戶顯式指定Storageformat_描述實(shí)際執(zhí)行時(shí)的內(nèi)存布局編譯過(guò)程中推導(dǎo)得到如果只用一個(gè) Format 字段會(huì)導(dǎo)致以下問(wèn)題語(yǔ)義丟失當(dāng) FE融合引擎將 NCHW 轉(zhuǎn)為 NC1HWC0 后原始的 NCHW 語(yǔ)義無(wú)處保存后續(xù)優(yōu)化 Pass 無(wú)法判斷這個(gè)張量原本代表什么優(yōu)化受限D(zhuǎn)ata Dump、Profiling 等調(diào)試場(chǎng)景需要將 NC1HWC0 的數(shù)據(jù)轉(zhuǎn)回 NCHW 供用戶理解沒(méi)有 Origin Format 就無(wú)法正確還原格式傳播混亂整網(wǎng)格式推導(dǎo)需要錨點(diǎn)如 Conv2D 的 data_format 屬性如果執(zhí)行格式覆蓋了語(yǔ)義格式推導(dǎo)的起點(diǎn)就丟失了具體而言以一個(gè) NCHW 張量[8, 3, 224, 224]為例字段值含義OriginFormatNCHW用戶定義的語(yǔ)義格式OriginShape[8, 3, 224, 224]用戶理解的維度StorageFormatNC1HWC0實(shí)際內(nèi)存布局StorageShape[8, 1, 224, 224, 16]C3 向上對(duì)齊到 16 后的實(shí)際存儲(chǔ)形態(tài)僅從 StorageShape[8, 1, 224, 224, 16]無(wú)法唯一還原其 OriginFormat ——它可能來(lái)自 NCHWC3也可能來(lái)自 NHWC因此兩個(gè)字段必須共存。2. 用戶使用場(chǎng)景2.1 離線編譯場(chǎng)景atc用戶使用 atc 工具將 ONNX/PB 模型編譯為 OM 文件時(shí)GE 需要自動(dòng)完成整網(wǎng)的格式推導(dǎo)模型輸入以 NCHW 或 NHWC 格式定義GE 推導(dǎo)出每個(gè)算子的 OriginFormatFE 根據(jù)算子能力和性能偏好選擇 StorageFormat最終生成包含正確內(nèi)存布局信息的 OM 模型2.2 在線訓(xùn)練/推理場(chǎng)景通過(guò) TorchAir 或 TFA 集成時(shí)框架傳入的計(jì)算圖可能未顯式標(biāo)注所有算子的格式。GE 需要在圖編譯階段推導(dǎo)圖中所有張量的 OriginFormat在PrepareRunningFormatRefiner階段根據(jù)用戶設(shè)置的storage_format屬性刷新 Data/NetOutput 節(jié)點(diǎn)確保 InferShape 在正確的格式上下文中執(zhí)行2.3 Data Dump 與 Profiling用戶在調(diào)試時(shí)需要查看算子的輸入輸出數(shù)據(jù)。數(shù)據(jù)在設(shè)備上以 StorageFormat 存儲(chǔ)但用戶理解的是 OriginFormat。GE 通過(guò)保存ATTR_NAME_DATA_DUMP_ORIGIN_FORMAT屬性在 Dump 時(shí)將數(shù)據(jù)從 StorageFormat 轉(zhuǎn)回 OriginFormat。3. 對(duì)外接口3.1 GeTensorDesc 上的格式接口GeTensorDesc定義于inc/graph_metadef/graph/ge_tensor.h是張量描述的核心類(lèi)提供以下格式相關(guān)接口GeTensorDesc ├── GetFormat() / SetFormat() // StorageFormat 的讀寫(xiě) ├── GetOriginFormat() / SetOriginFormat() // OriginFormat 的讀寫(xiě) ├── GetShape() / SetShape() // StorageShape 的讀寫(xiě) ├── GetOriginShape() / SetOriginShape() // OriginShape 的讀寫(xiě)在graph_metadef/graph/normal_graph/tensor.cc的TensorDescImpl類(lèi)中可以看到這兩個(gè)字段獨(dú)立存儲(chǔ)format_對(duì)應(yīng) StorageFormat默認(rèn)FORMAT_NDorigin_format_對(duì)應(yīng) OriginFormat默認(rèn)FORMAT_NDorigin_format_is_set_標(biāo)記 OriginFormat 是否被顯式設(shè)置3.2 StorageFormat 描述體在運(yùn)行時(shí)gert 命名空間StorageFormat是一個(gè)同時(shí)攜帶 Origin 和 Storage 信息的描述體定義于inc/graph_metadef/external/graph/types.h其構(gòu)造方式為StorageFormat(origin_format, storage_format, expand_dims_type)在graph_metadef/register/shape_inference.cc的GetTensorHolder函數(shù)中可以看到創(chuàng)建 Tensor 時(shí)會(huì)將GeTensorDesc的兩個(gè)格式字段分別映射{input_desc.GetOriginFormat(), input_desc.GetFormat(), {}}即StorageFormat描述體的第一個(gè)參數(shù)是 OriginFormat第二個(gè)參數(shù)是 StorageFormat。3.3 算子 InferFormat 注冊(cè)接口算子開(kāi)發(fā)者可通過(guò)以下方式注冊(cè)格式推導(dǎo)函數(shù)IMPL_OP(OpType).InferFormat(infer_format_func)其中infer_format_func簽名為UINT32(InferFormatContext *context)。V2 接口通過(guò)InferFormatContext提供更結(jié)構(gòu)化的輸入輸出訪問(wèn)。在graph_metadef/register/shape_inference.cc的UpdateOpDescOutFormat函數(shù)中V2 推導(dǎo)完成后會(huì)將結(jié)果寫(xiě)回 OpDescdesc-SetOriginFormat(format-GetOriginFormat()); desc-SetFormat(format-GetStorageFormat());3.4 用戶指定 StorageFormat用戶可通過(guò)在 TensorDesc 上設(shè)置屬性來(lái)指定算子的 StorageFormatAttrUtils::SetInt(tensor_desc, ATTR_NAME_STORAGE_FORMAT, format_value); AttrUtils::SetListInt(tensor_desc, ATTR_NAME_STORAGE_SHAPE, shape_dims);這些屬性在compiler/graph/preprocess/graph_prepare.cc的PrepareRunningFormatRefiner階段被消費(fèi)用于刷新 Data 節(jié)點(diǎn)和 NetOutput 節(jié)點(diǎn)的格式。3.5 Format 枚舉定義GE 支持的格式類(lèi)型定義于inc/framework/executor_c/types.h核心格式包括格式說(shuō)明典型用途FORMAT_NDN 維張量默認(rèn)格式不攜帶特殊語(yǔ)義FORMAT_NCHWN-C-H-W用戶側(cè)常見(jiàn)的卷積格式FORMAT_NHWCN-H-W-CTensorFlow 默認(rèn)格式FORMAT_NC1HWC0N-C1-H-W-C0Conv2D 在昇騰上的親和格式FORMAT_FRACTAL_ZC1HW-N1-N0-C0Conv2D Filter 的親和格式FORMAT_FRACTAL_NZN1-N0-C1-C0MatMul 權(quán)重的親和格式4. 具體實(shí)現(xiàn)4.1 整體流程格式推導(dǎo)在 GE 編譯流程中的位置如下關(guān)鍵順序先推導(dǎo) OriginFormat再做 InferShape最后處理 StorageFormat。這是因?yàn)?InferShape 需要在正確的 OriginFormat 上下文中執(zhí)行Shape 的維度含義依賴(lài) Format而 StorageFormat 的選擇發(fā)生在 FE 階段。4.2 Origin Format 推導(dǎo)FormatRefinerOrigin Format 推導(dǎo)的核心實(shí)現(xiàn)在graph_metadef/graph/refiner/format_refiner.cc的FormatRefiner::InferOrigineFormat函數(shù)中。4.2.1 推導(dǎo)算法推導(dǎo)采用錨點(diǎn)擴(kuò)散策略具體步驟錨點(diǎn)識(shí)別GetAnchorPoints遍歷圖中所有節(jié)點(diǎn)找到輸入/輸出中存在非 ND 格式的節(jié)點(diǎn)作為錨點(diǎn)。這些節(jié)點(diǎn)通常是 Conv2D、Pooling 等對(duì)格式敏感的算子它們通過(guò)屬性如data_format已攜帶格式信息。錨點(diǎn)刷新RefreshOriginFormatOfAnchor對(duì)錨點(diǎn)節(jié)點(diǎn)如果origin_format仍為 ND 或 RESERVED則將其format值復(fù)制到origin_format。這確保錨點(diǎn)自身的 OriginFormat 被正確建立。雙向擴(kuò)散AnchorProcess向后推導(dǎo)BackInferProcess從錨點(diǎn)的輸入端出發(fā)沿?cái)?shù)據(jù)流反向傳播格式。對(duì)每個(gè)上游節(jié)點(diǎn)如果其origin_format為 ND 且未鎖定則將錨點(diǎn)的格式傳遞過(guò)去。向前推導(dǎo)ForwardInferProcess從錨點(diǎn)的輸出端出發(fā)沿?cái)?shù)據(jù)流正向傳播格式。Data 節(jié)點(diǎn)兜底DataNodeFormatProcess對(duì)于推導(dǎo)過(guò)程中未被觸及的 Data 節(jié)點(diǎn)通常是因?yàn)槿鄙俑袷藉^點(diǎn)使用圖的全局data_format參數(shù)統(tǒng)一設(shè)置格式。4.2.2 格式傳播的中斷條件推導(dǎo)過(guò)程在以下情況會(huì)中斷節(jié)點(diǎn)格式已鎖定ATTR_NAME_FORMAT_LOCKED為 true某些算子的格式不應(yīng)被推導(dǎo)過(guò)程覆蓋節(jié)點(diǎn)類(lèi)型為維度變化算子PERMUTE、EXPANDDIMS、SQUEEZE當(dāng)維度數(shù)小于 4 時(shí)維度語(yǔ)義不確定不應(yīng)傳播格式遇到標(biāo)量dim_num 0標(biāo)量無(wú)格式語(yǔ)義遇到 ND 格式ND 表示無(wú)特定格式語(yǔ)義傳播到此處停止遇到 NetOutput作為圖的邊界不繼續(xù)向前傳播4.2.3 Ref 反射機(jī)制對(duì)于 If/Case 等控制流算子GE 通過(guò)RefRelations建立子圖與主圖之間 Data 節(jié)點(diǎn)的反射關(guān)系。當(dāng)主圖中某個(gè) Data 的格式被推導(dǎo)后通過(guò)ReflectionProcess將格式同步到子圖對(duì)應(yīng)的 Data 節(jié)點(diǎn)及其父節(jié)點(diǎn)的輸入。4.2.4 算子自定義推導(dǎo)除了默認(rèn)的格式傳播機(jī)制算子還可注冊(cè)自定義的 InferFormat 函數(shù)。在NodeUtilsEx::InferOriginFormat中會(huì)調(diào)用OpDescUtilsEx::CallInferFormatFunc優(yōu)先使用算子注冊(cè)的推導(dǎo)函數(shù)否則使用DefaultInferFormat將第一個(gè)非 ND 格式傳播到所有輸入輸出。4.3 Storage Format 的確定Storage Format 的確定分為兩條路徑4.3.1 用戶顯式指定路徑在compiler/graph/preprocess/graph_prepare.cc的UpdateDataNetOutputByStorageFormat函數(shù)中對(duì) Data 節(jié)點(diǎn)從ATTR_NAME_STORAGE_FORMAT和ATTR_NAME_STORAGE_SHAPE屬性讀取用戶指定的 StorageFormat調(diào)用ModifyTensorDescStorageFormatAndShape刷新 TensorDesc對(duì) NetOutput 節(jié)點(diǎn)同樣讀取屬性并刷新確保輸出格式正確對(duì) ConstPlaceHolder 節(jié)點(diǎn)處理常量的存儲(chǔ)格式ModifyTensorDescStorageFormatAndShape函數(shù)的核心操作根據(jù) StorageFormat 和 OriginFormat 計(jì)算存儲(chǔ)形態(tài)StorageShape包括維度的擴(kuò)展和對(duì)齊調(diào)用SetFormat設(shè)置 StorageFormat注意不是SetOriginFormat調(diào)用SetShape設(shè)置 StorageShape計(jì)算并設(shè)置 Tensor 的內(nèi)存大小4.3.2 FE 自動(dòng)選擇路徑對(duì)于計(jì)算密集型算子如 Conv2D、MatMulFE融合引擎在算子編譯階段根據(jù)算子的能力選擇最優(yōu)的 StorageFormat。這一過(guò)程由compiler/engines/nn_engine/optimizer/format_selector/下的 FormatSelector 體系實(shí)現(xiàn)FormatDtypeOpBuiltinSelector內(nèi)置算子的格式選擇FormatDtypeOpKernelSelector基于算子內(nèi)核信息的格式選擇FormatDtypeOpCustomizeSelector自定義算子的格式選擇這些 Selector 通過(guò)FormatDtypeManagerBase協(xié)調(diào)最終由FormatDtypeSetter將選擇的格式設(shè)置到圖節(jié)點(diǎn)上。4.4 運(yùn)行時(shí)格式推導(dǎo)InferShape 階段在運(yùn)行時(shí)的 InferShape 階段格式信息通過(guò)InferFormatContextV2 接口傳遞給算子的推導(dǎo)函數(shù)。在graph_metadef/register/shape_inference.cc的InferFormatOnCompile函數(shù)中構(gòu)造InferFormatContext為每個(gè)輸入/輸出創(chuàng)建CompileTimeTensorDesc其中同時(shí)包含origin_format和storage_format調(diào)用算子注冊(cè)的infer_format_func將推導(dǎo)結(jié)果通過(guò)UpdateOpDescOutFormat寫(xiě)回 OpDesc在GetTensorHolder函數(shù)中可以看到 Tensor 的創(chuàng)建方式gert::Tensor(storage_shape, {input_desc.GetOriginFormat(), input_desc.GetFormat(), {}}, input_desc.GetDataType())其中StorageShape描述體同時(shí)攜帶 Origin 和 Storage 的 Shape 信息。4.5 格式差異檢測(cè)與 TransData 插入在runtime/v2/graph_builder/storage_format.cc中DiffStorageFormat函數(shù)檢測(cè)一個(gè)張量的 OriginFormat 與 StorageFormat 是否不同或 Shape 是否不同return td-GetFormat() ! td-GetOriginFormat() || td-GetShape().GetDims() ! td-GetOriginShape().GetDims()AnyDiffStorageFormat檢查一個(gè)節(jié)點(diǎn)的所有輸入輸出中是否存在任何格式差異。如果存在差異則在 Lowering 階段插入 TransData 算子來(lái)完成實(shí)際的格式轉(zhuǎn)換。4.6 InferShape 中的格式刷新在runtime/v2/kernel/common_kernel_impl/infer_shape_compatible.cc的兼容性 InferShape 中有一個(gè)關(guān)鍵的格式處理// RT1時(shí)算子的infershape只能拿到format字段但是卻需要用origin format input_desc-SetFormat(input_desc_in_context-GetOriginFormat()); input_desc-SetOriginFormat(input_desc_in_context-GetOriginFormat());這表明在 RT1運(yùn)行時(shí)第一版兼容模式下InferShape 只能獲取到format字段但實(shí)際需要的是 OriginFormat。因此需要從 Context 中正確取出 OriginFormat 并同步設(shè)置到format和origin_format兩個(gè)字段。在runtime/v2/kernel/common_kernel_impl/infer_shape.h的TransformOutputShape函數(shù)中當(dāng) OriginFormat 與 StorageFormat 不同時(shí)會(huì)調(diào)用ShapeTransferAccordingToFormat::TransferShape將 OriginShape 轉(zhuǎn)換為 StorageShapeif (output_td-GetOriginFormat() output_td-GetStorageFormat()) { // 格式相同無(wú)需轉(zhuǎn)換 return GRAPH_SUCCESS; } // 格式不同需要根據(jù) StorageFormat 計(jì)算 StorageShape TransferShape(origin_format, storage_format, data_type, storage_shape)4.7 整網(wǎng)編譯流程中的調(diào)用時(shí)機(jī)在compiler/graph/manager/graph_manager.cc中格式相關(guān)的階段按以下順序執(zhí)行PrepareRunningFormatRefiner ← StorageFormat 刷新 → UpdateDataNetOutputByStorageFormat → VariablePrepareOpPass → UpdateInputOutputByOptions → UpdateVariableFormats在此之前的 GraphPrepare::GenerateInfershapeGraph 中InferOriginFormat ← OriginFormat 推導(dǎo) → FormatRefiner::InferOrigineFormat4.8 格式優(yōu)化 Pass在 OriginFormat 推導(dǎo)和 StorageFormat 確定之后compiler/graph/passes/format_optimize/下的多個(gè) Pass 負(fù)責(zé)優(yōu)化 TransData 的插入和消除Pass功能TransOpSymmetryEliminationPass消除對(duì)稱(chēng)的格式轉(zhuǎn)換對(duì)如 NCHW→NC1HWC0→NCHWTransOpBreadthFusionPass將同一節(jié)點(diǎn)的多個(gè)輸出側(cè) TransData 合并TransOpWithoutReshapeFusionPass融合不含 Reshape 的連續(xù) TransDataTransposeTransDataPass將 Transpose 與 TransData 合并優(yōu)化UnchangedTransposeRemovePass移除不改變數(shù)據(jù)的 TransposeCastRemovePass移除不必要的 Cast其中TransOpWithoutReshapeFusionPass只處理 shape、format 和轉(zhuǎn)換算子輸入 dtype 均連續(xù)的轉(zhuǎn)換鏈如果轉(zhuǎn)換算子輸入 dtype 與上游輸出 dtype 不一致則保留原鏈路避免誤刪轉(zhuǎn)換節(jié)點(diǎn)。5. 關(guān)鍵設(shè)計(jì)決策5.1 為什么 FormatRefiner 使用錨點(diǎn)擴(kuò)散而非全圖遍歷錨點(diǎn)擴(kuò)散的優(yōu)勢(shì)在于避免無(wú)意義傳播大量 ElementWise 算子如 Add、ReLU對(duì)格式不敏感其格式應(yīng)與上游保持一致不需要單獨(dú)推導(dǎo)降低復(fù)雜度只在格式語(yǔ)義發(fā)生變化的節(jié)點(diǎn)錨點(diǎn)附近做推導(dǎo)而非遍歷全圖支持控制流錨點(diǎn)擴(kuò)散結(jié)合 RefRelations 可以自然處理子圖間的格式同步5.2 為什么 OriginFormat 和 StorageFormat 在 GeTensorDesc 中都存儲(chǔ)為format_和origin_format_這種設(shè)計(jì)的核心考慮是origin_format_是 FormatRefiner 階段確定的表達(dá)用戶語(yǔ)義一旦確定不應(yīng)被修改format_初始與origin_format_相同后續(xù)在 FE 階段被修改為 StorageFormat兩個(gè)字段共存允許任何時(shí)刻對(duì)比它們來(lái)判斷是否需要 TransData5.3 為什么 StorageFormat 描述體同時(shí)攜帶 Origin 信息StorageFormatgert 命名空間雖然名字叫Storage但實(shí)際上是一個(gè)復(fù)合描述體。原因是僅靠 StorageFormat 的值如 NC1HWC0無(wú)法唯一還原其 OriginFormat可能是 NCHW 或 NHWC必須同時(shí)保存兩者。這使得運(yùn)行時(shí)無(wú)需回溯到 OpDesc 就能獲得完整的格式上下文。6. 數(shù)據(jù)流總結(jié)以下是一個(gè)典型 Conv2D 網(wǎng)絡(luò)中格式信息的完整生命周期【免費(fèi)下載鏈接】geGEGraph Engine是面向昇騰的圖編譯器和執(zhí)行器提供了計(jì)算圖優(yōu)化、多流并行、內(nèi)存復(fù)用和模型下沉等技術(shù)手段加速模型執(zhí)行效率減少模型內(nèi)存占用。 GE 提供對(duì) PyTorch、TensorFlow 前端的友好接入能力并同時(shí)支持 onnx、pb 等主流模型格式的解析與編譯。項(xiàng)目地址: https://gitcode.com/cann/ge創(chuàng)作聲明:本文部分內(nèi)容由AI輔助生成(AIGC),僅供參考