load("@llvm-project//mlir:tblgen.bzl", "gentbl_cc_library", "td_library") # Copyright 2023 The JAX Authors. # # Licensed under the Apache License, Version 2.0 (the "License"); # you may not use this file except in compliance with the License. # You may obtain a copy of the License at # # https://www.apache.org/licenses/LICENSE-2.0 # # Unless required by applicable law or agreed to in writing, software # distributed under the License is distributed on an "AS IS" BASIS, # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # See the License for the specific language governing permissions and # limitations under the License. load("@rules_cc//cc:cc_library.bzl", "cc_library") licenses(["notice"]) package( default_applicable_licenses = [], default_visibility = [ "//visibility:public", ], ) ################################################################################ # TPU dialect cc_library( name = "tpu_dialect", srcs = [ "jaxlib/mosaic/dialect/tpu/array_util.cc", "jaxlib/mosaic/dialect/tpu/layout.cc", "jaxlib/mosaic/dialect/tpu/tpu_dialect.cc", "jaxlib/mosaic/dialect/tpu/tpu_ops.cc", "jaxlib/mosaic/dialect/tpu/util.cc", "jaxlib/mosaic/dialect/tpu/vreg_util.cc", ], hdrs = [ "jaxlib/mosaic/dialect/tpu/array_util.h", "jaxlib/mosaic/dialect/tpu/layout.h", "jaxlib/mosaic/dialect/tpu/tpu_dialect.h", "jaxlib/mosaic/dialect/tpu/util.h", "jaxlib/mosaic/dialect/tpu/vreg_util.h", ], # compatible with libtpu deps = [ ":stringify_util", ":tpu_inc_gen", "@com_google_absl//absl/algorithm:container", "@com_google_absl//absl/hash", "@com_google_absl//absl/log:check", "@com_google_absl//absl/status", "@com_google_absl//absl/status:statusor", "@com_google_absl//absl/strings", "@com_google_absl//absl/strings:str_format", "@com_google_absl//absl/types:span", "@llvm-project//llvm:Support", "@llvm-project//mlir:ArithDialect", "@llvm-project//mlir:CommonFolders", "@llvm-project//mlir:DialectUtils", "@llvm-project//mlir:FuncDialect", "@llvm-project//mlir:IR", "@llvm-project//mlir:MathDialect", "@llvm-project//mlir:MemRefDialect", "@llvm-project//mlir:Pass", "@llvm-project//mlir:SCFDialect", "@llvm-project//mlir:Support", "@llvm-project//mlir:VectorDialect", "@xla//xla:array", "@xla//xla:shape_util", "@xla//xla/tsl/platform:statusor", ], ) cc_library( name = "tpu_serde_pass", srcs = ["jaxlib/mosaic/dialect/tpu/transforms/serde.cc"], hdrs = ["jaxlib/mosaic/dialect/tpu/transforms/serde.h"], # compatible with libtpu deps = [ ":pass_boilerplate", ":serde", ":tpu_dialect", "@llvm-project//llvm:Support", "@llvm-project//mlir:ArithDialect", "@llvm-project//mlir:DataLayoutInterfaces", "@llvm-project//mlir:IR", "@llvm-project//mlir:Pass", "@llvm-project//mlir:Support", "@llvm-project//mlir:VectorDialect", ], ) gentbl_cc_library( name = "tpu_inc_gen", # compatible with libtpu tbl_outs = { "jaxlib/mosaic/dialect/tpu/tpu_ops.h.inc": [ "-gen-op-decls", "-dialect=tpu", ], "jaxlib/mosaic/dialect/tpu/tpu_ops.cc.inc": [ "-gen-op-defs", "-dialect=tpu", ], "jaxlib/mosaic/dialect/tpu/tpu_dialect.h.inc": [ "-gen-dialect-decls", "-dialect=tpu", ], "jaxlib/mosaic/dialect/tpu/tpu_dialect.cc.inc": [ "-gen-dialect-defs", "-dialect=tpu", ], "jaxlib/mosaic/dialect/tpu/tpu_enums.h.inc": [ "-gen-enum-decls", "-dialect=tpu", ], "jaxlib/mosaic/dialect/tpu/tpu_enums.cc.inc": [ "-gen-enum-defs", "-dialect=tpu", ], "jaxlib/mosaic/dialect/tpu/tpu_attr_defs.h.inc": [ "-gen-attrdef-decls", "-dialect=tpu", "--attrdefs-dialect=tpu", ], "jaxlib/mosaic/dialect/tpu/tpu_attr_defs.cc.inc": [ "-gen-attrdef-defs", "-dialect=tpu", "--attrdefs-dialect=tpu", ], "jaxlib/mosaic/dialect/tpu/tpu_type_defs.h.inc": [ "-gen-typedef-decls", "-dialect=tpu", "--typedefs-dialect=tpu", ], "jaxlib/mosaic/dialect/tpu/tpu_type_defs.cc.inc": [ "-gen-typedef-defs", "-dialect=tpu", "--typedefs-dialect=tpu", ], "jaxlib/mosaic/dialect/tpu/tpu_passes.h.inc": [ "-gen-pass-decls", "-name=TPU", ], "jaxlib/mosaic/dialect/tpu/integrations/c/tpu_passes.capi.h.inc": [ "-gen-pass-capi-header", "--prefix=TPU", ], "jaxlib/mosaic/dialect/tpu/integrations/c/tpu_passes.capi.cc.inc": [ "-gen-pass-capi-impl", "--prefix=TPU", ], }, tblgen = "@llvm-project//mlir:mlir-tblgen", td_file = "jaxlib/mosaic/dialect/tpu/tpu_ops.td", deps = [":tpu_ops_td_files"], ) td_library( name = "tpu_td_files", srcs = [ "jaxlib/mosaic/dialect/tpu/tpu.td", ], # compatible with libtpu deps = [ "@llvm-project//mlir:BuiltinDialectTdFiles", ], ) td_library( name = "tpu_ops_td_files", srcs = [ "jaxlib/mosaic/dialect/tpu/tpu_ops.td", ], # compatible with libtpu deps = [ ":tpu_td_files", "@llvm-project//mlir:BuiltinDialectTdFiles", "@llvm-project//mlir:ControlFlowInterfacesTdFiles", "@llvm-project//mlir:InferTypeOpInterfaceTdFiles", "@llvm-project//mlir:OpBaseTdFiles", "@llvm-project//mlir:PassBaseTdFiles", "@llvm-project//mlir:SideEffectInterfacesTdFiles", ], ) # C API targets TPU_CAPI_SOURCES = [ "jaxlib/mosaic/dialect/tpu/integrations/c/tpu_dialect.cc", "jaxlib/mosaic/dialect/tpu/integrations/c/tpu_passes.capi.cc.inc", ] TPU_CAPI_HEADERS = [ "jaxlib/mosaic/dialect/tpu/integrations/c/tpu_dialect.h", "jaxlib/mosaic/dialect/tpu/integrations/c/tpu_passes.capi.h.inc", ] cc_library( name = "tpu_dialect_capi", srcs = TPU_CAPI_SOURCES, hdrs = TPU_CAPI_HEADERS, deps = [ ":tpu_dialect", ":tpu_inc_gen", ":tpu_serde_pass", "@llvm-project//llvm:Support", "@llvm-project//mlir:CAPIIR", "@llvm-project//mlir:FuncDialect", "@llvm-project//mlir:IR", "@llvm-project//mlir:Support", ], ) # Header-only target, used when using the C API from a separate shared library. cc_library( name = "tpu_dialect_capi_headers", hdrs = TPU_CAPI_HEADERS, deps = [ ":tpu_inc_gen", "@llvm-project//mlir:CAPIIRHeaders", ], ) # Alwayslink target, used when exporting the C API from a shared library. cc_library( name = "tpu_dialect_capi_objects", srcs = TPU_CAPI_SOURCES, hdrs = TPU_CAPI_HEADERS, deps = [ ":tpu_dialect", ":tpu_inc_gen", ":tpu_serde_pass", "@llvm-project//llvm:Support", "@llvm-project//mlir:CAPIIR", "@llvm-project//mlir:FuncDialect", "@llvm-project//mlir:IR", "@llvm-project//mlir:Support", ], alwayslink = True, ) cc_library( name = "pass_boilerplate", hdrs = ["jaxlib/mosaic/pass_boilerplate.h"], # compatible with libtpu deps = [ "@llvm-project//mlir:IR", "@llvm-project//mlir:Pass", "@llvm-project//mlir:Support", ], ) cc_library( name = "serde", srcs = ["jaxlib/mosaic/serde.cc"], hdrs = ["jaxlib/mosaic/serde.h"], # compatible with libtpu deps = [ "@llvm-project//llvm:Support", "@llvm-project//mlir:DataLayoutInterfaces", "@llvm-project//mlir:IR", "@llvm-project//mlir:Support", ], ) cc_library( name = "stringify_util", hdrs = ["jaxlib/mosaic/dialect/tpu/stringify_util.h"], # compatible with libtpu deps = ["@llvm-project//llvm:Support"], )