Mecanismo de Cache de Compilação de Subgrafos no CINN

Entrada da Compilação de Subgrafos

O processo de compilação de subgrafos no CINN inicia-se através da iteração sobre clusters identificados:

for (const auto& cluster_nodes : graph_clusters) {
    NodeCollection cluster_elements(cluster_nodes.begin(), cluster_nodes.end());
    
    NodeCollection input_nodes, output_nodes, internal_nodes;
    ClassifyClusterNodes(cluster_elements,
                        excluded_vars,
                        &input_nodes,
                        &output_nodes,
                        &internal_nodes,
                        is_inference_phase,
                        preserved_variables);

    auto subgraph = BuildSubGraph(
        cluster_elements, internal_nodes, input_nodes, output_nodes);

    auto compilation_id = cinn_compiler->RegisterGraph(std::move(subgraph));
    
    ReplaceClusterWithCinnOp(cluster_elements,
                            input_nodes,
                            output_nodes,
                            internal_nodes,
                            compilation_id,
                            main_graph);
}

Reigstro de Grafos no Compilador

A função RegisterGraph armazena o grafo para compilação futura:

int64_t CinnCompiler::RegisterGraph(std::unique_ptr<graph> graph) {
    int64_t graph_id = std::hash<graph>()((&(*graph)));
    stored_graphs_[graph_id] = std::move(graph);
    return graph_id;
}</graph></graph>

Mecanismo de Hash para Cache

O sistema utiliza múltiplas estratégias de hash para identificação de cache:

size_t StructureCacheKey::ComputeGraphHash(const ir::Graph& graph) {
    auto node_comparator = [](ir::Node* a, ir::Node* b) {
        return (a->Name() == b->Name()) ? (a->id() < b->id())
                                      : (a->Name() < b->Name());
    };

    std::set<:node bool="" ir::node=""> sorted_nodes(node_comparator),
                                                       sorted_outputs(node_comparator);
    sorted_nodes.insert(graph.Nodes().begin(), graph.Nodes().end());

    std::string hash_string;
    for (ir::Node* node : sorted_nodes) {
        hash_string.append(node->Name());
        sorted_outputs.clear();
        sorted_outputs.insert(node->outputs.begin(), node->outputs.end());
        for (auto* output : sorted_outputs) {
            hash_string.append(output->Name());
        }
    }

    return std::hash<:string>()(hash_string);
}</:string></:node>

Exemplo de hash gerado para um modelo Bert:

cumsumcumsum_0.tmp_0cumsum_0.tmp_0elementwise_subelementwise_subtmp_0feedinput_idsfetchfill_any_likefull_like_0.tmp_0full_like_0.tmp_0cumsumelementwise_subinput_idsfill_any_liketmp_0fetch

Disparo da Compilação

A compilação é acionada durante a execução do kernel de lançamento:

template <typename devicecontext="" t="" typename="">
class CinnLaunchOpKernel : public framework::OpKernel<t> {
public:
    void Compute(const framework::ExecutionContext& ctx) const override {
        const auto& compile_key = ctx.template Attr<int64_t>(kCompilationKey);
        const auto& compiled_obj = CinnCompiler::GetInstance()->Compile(
            compile_key, input_tensors, target, stream);
    }
};</int64_t></t></typename>

Resolução de Shape Dinâmico

A inferência de shapes é realizada durante a simbolização do grafo:

void CinnGraphSymbolization::ExecuteOp(const CinnOpDesc& op_desc,
                                      const OpMapperContext& ctx) const {
    const auto& op_type = op_desc.Type();
    auto* kernel = ::cinn::frontend::OpMapperRegistry::Global()->Find(op_type);
    kernel->Execute(op_desc, ctx);
}

Implementação do Cache de Compilação

O sistema emprega cache duplo baseado em endereço e estrutura:

const CinnCompiledObject& CinnCompiler::Compile(
    const Graph& graph,
    const std::map<:string const="" phi::densetensor="">& input_tensors,
    const Target& target,
    void* stream) {
    
    AddressCacheKey current_addr_key(graph, input_tensors, target.arch_str());
    StructureCacheKey current_struct_key;

    if (!address_cache_.count(current_addr_key)) {
        current_struct_key.GenerateKey(graph, input_tensors, target.arch_str());
        if (!structure_cache_.count(current_struct_key)) {
            std::int64_t compile_count = compilation_counter_.fetch_add(1);
            auto compilation_result =
                CompileGraphInternal(graph, input_tensors, target, compile_count, stream);
            
            std::unique_lock<:mutex> lock(mutex_);
            if (!structure_cache_.count(current_struct_key)) {
                structure_cache_[current_struct_key] = compile_count;
                cache_index_.emplace(compile_count, std::move(compilation_result));
            }
            if (!address_cache_.count(current_addr_key)) {
                address_cache_[current_addr_key] = structure_cache_.at(current_struct_key);
            }
        } else {
            std::unique_lock<:mutex> lock(mutex_);
            if (!address_cache_.count(current_addr_key)) {
                address_cache_[current_addr_key] = structure_cache_.at(current_struct_key);
            }
        }
    }
    return *cache_index_.at(address_cache_.at(current_addr_key));
}</:mutex></:mutex></:string>

Compilação do Grafo

O processo principal de compilação:

std::unique_ptr<cinncompiledobject> CinnCompiler::CompileGraphInternal(
    const ir::Graph& graph,
    const std::map<:string const="" phi::densetensor="">& input_tensors,
    const Target& target,
    std::int64_t compilation_id,
    void* stream) const {
    
    CinnGraphSymbolization symbolizer{compilation_id, graph, target, input_tensors};
    auto frontend_program = symbolizer();
    auto output_var_ids = symbolizer.GetOutputVarIds();

    auto optimized_graph = ApplyOptimizations(&frontend_program, output_var_ids, target);
    
    auto scope = CreateScope(target, optimized_graph);
    auto graph_compiler =
        std::make_unique<graphcompiler>(target, scope, optimized_graph);
    
    GraphCompiler::CompilationOptions opts;
    opts.instantiate_variables = false;
    
    auto compilation_output =
        graph_compiler->Build(opts, std::move(output_var_ids), stream);
    
    auto compiled_object = std::make_unique<cinncompiledobject>();
    *compiled_object = {std::move(graph_compiler),
                       nullptr, // auto_tuner
                       std::move(compilation_output.runtime_program),
                       scope,
                       symbolizer.GetVarMapping()};
    compiled_object->cache_index = compilation_id;
    compiled_object->execution_context =
        std::make_unique<:details::cinnexecutioncontext>(graph,
                                                                  *compiled_object);
    ValidateCompilation(graph, input_tensors, *compiled_object);
    return compiled_object;
}</:details::cinnexecutioncontext></cinncompiledobject></graphcompiler></:string></cinncompiledobject>

Comparação com TVM

No ecossistema TVM, o TECompilerImpl desempenha papel similar:

class TECompilerImpl : public TECompilerNode {
public:
    CachedFunc Lower(const CCacheKey& key) {
        return LowerInternal(key, global_var_provider)->cached_func;
    }
    
    PackedFunc JIT(const CCacheKey& key) final {
        CCacheValue result = LowerInternal(key, GlobalVarSupplier(NameSupplier("")));
        if (result->packed_func != nullptr) {
            return result->packed_func;
        }
        auto module = Build(result->cached_func->functions, key->target, Target(nullptr));
        result->packed_func = module.GetFunction(result->cached_func->primary_func_var->name_hint);
        return result->packed_func;
    }
    
private:
    std::unordered_map<ccachekey ccachevalue=""> compilation_cache_;
    std::unordered_map<ccachekey ccachevalue=""> shape_function_cache_;
};</ccachekey></ccachekey>

O prcoesso de build no TVM:

runtime::Module Build(const Map<target irmodule="">& inputs, const Target& host_target) {
    return ConvertTIRToRuntime(inputs, host_target);
}

runtime::Module ConvertTIRToRuntime(const Map<target irmodule="">& input_modules,
                                   const Target& host_target) {
    std::vector<:module> device_modules;
    // Processamento de módulos host e device
    // Compilação separada e integração
    runtime::Module host_module = codegen::Build(host_module_all, host_target);
    for (const auto& device_module : device_modules) {
        if (device_module.operator->()) {
            host_module.Import(device_module);
        }
    }
    return host_module;
}</:module></target></target>

Tags: CINN compilacao subgrafo cache TVM

Publicado em 8-15 07:41