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