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>