diff --git a/qa/L0_pytorch_lint/check_torch_boundary.py b/qa/L0_pytorch_lint/check_torch_boundary.py new file mode 100644 index 0000000000..4a0ad47a45 --- /dev/null +++ b/qa/L0_pytorch_lint/check_torch_boundary.py @@ -0,0 +1,146 @@ +#!/usr/bin/env python3 +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +"""Enforce the single TE<->PyTorch binary boundary. + +Only ``transformer_engine/pytorch/csrc/torch_backend.{h,cpp}`` may include +libtorch/ATen/c10 headers or name ``at::`` / ``c10::`` / ``c10d::`` / +``torch::`` symbols. Every other translation unit must talk to PyTorch through +the aliases and free functions that ``torch_backend.h`` exposes. + +Files that have not been migrated yet are grandfathered in via +``torch_boundary_allowlist.txt`` (paths relative to the csrc dir). The guard +fails if: + * a non-boundary, non-allowlisted file touches the torch ABI, or + * an allowlisted file no longer touches it (stale entry -> remove it), or + * an allowlisted path does not exist. + +Comments and string literals are stripped before scanning, so mentioning the +tokens in a comment is fine. + +Usage: check_torch_boundary.py [TE_ROOT] (default: $TE_PATH or cwd) +""" + +import os +import re +import sys +from pathlib import Path + +CSRC_REL = "transformer_engine/pytorch/csrc" +BOUNDARY = {"torch_backend.h", "torch_backend.cpp"} +SOURCE_SUFFIXES = {".cpp", ".h", ".hpp", ".cuh", ".cu"} + +# Include of a libtorch/ATen/c10 header, or a qualified at::/c10::/c10d::/torch:: name. +INCLUDE_RE = re.compile(r'#\s*include\s*[<"](?:torch|ATen|c10)/') +SYMBOL_RE = re.compile(r'\b(?:at|c10|c10d|torch)::') + + +def strip_comments_and_strings(text: str) -> str: + """Blank out // and /* */ comments and "..."/'...' literals (newlines kept).""" + out = [] + i, n = 0, len(text) + while i < n: + c = text[i] + two = text[i:i + 2] + if two == "//": + i += 2 + while i < n and text[i] != "\n": + i += 1 + elif two == "/*": + i += 2 + while i < n and text[i:i + 2] != "*/": + out.append("\n" if text[i] == "\n" else " ") + i += 1 + i += 2 + elif c in "\"'": + quote = c + out.append(" ") + i += 1 + while i < n and text[i] != quote: + if text[i] == "\\" and i + 1 < n: + i += 1 + out.append("\n" if text[i] == "\n" else " ") + i += 1 + i += 1 + out.append(" ") + else: + out.append(c) + i += 1 + return "".join(out) + + +def scan(path: Path): + """Return list of (lineno, text) lines that touch the torch ABI.""" + raw = path.read_text(encoding="utf-8", errors="replace") + code = strip_comments_and_strings(raw) + hits = [] + for lineno, line in enumerate(code.splitlines(), start=1): + if INCLUDE_RE.search(line) or SYMBOL_RE.search(line): + hits.append(lineno) + return hits + + +def main() -> int: + root = Path(sys.argv[1] if len(sys.argv) > 1 else os.environ.get("TE_PATH", ".")).resolve() + csrc = root / CSRC_REL + if not csrc.is_dir(): + print(f"error: csrc dir not found: {csrc}", file=sys.stderr) + return 2 + + allowlist_file = root / "qa/L0_pytorch_lint/torch_boundary_allowlist.txt" + allowlist = set() + if allowlist_file.is_file(): + for line in allowlist_file.read_text().splitlines(): + line = line.split("#", 1)[0].strip() + if line: + allowlist.add(line) + + violations = [] # (rel, lineno) touching torch outside the boundary + still_allowlisted = set() # allowlisted files that still touch torch (expected) + + for path in sorted(csrc.rglob("*")): + if path.suffix not in SOURCE_SUFFIXES or not path.is_file(): + continue + rel = path.relative_to(csrc).as_posix() + if rel in BOUNDARY: + continue + hits = scan(path) + if not hits: + continue + if rel in allowlist: + still_allowlisted.add(rel) + else: + violations.extend((rel, ln) for ln in hits) + + stale = sorted(allowlist - still_allowlisted) + missing = sorted(p for p in allowlist if not (csrc / p).is_file()) + + ok = True + if violations: + ok = False + print("TE<->torch boundary violations (route these through torch_backend.h):") + for rel, ln in violations: + print(f" {CSRC_REL}/{rel}:{ln}") + if stale: + ok = False + print("\nStale allowlist entries (migrated -- remove from " + "torch_boundary_allowlist.txt):") + for rel in stale: + print(f" {rel}") + if missing: + ok = False + print("\nAllowlist entries pointing at non-existent files (remove them):") + for rel in missing: + print(f" {rel}") + + if ok: + print(f"torch boundary OK: only torch_backend.* touches the torch ABI " + f"({len(still_allowlisted)} file(s) still pending migration).") + return 0 + return 1 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/qa/L0_pytorch_lint/test.sh b/qa/L0_pytorch_lint/test.sh index f08dd8a03d..fdfe5fc9e9 100755 --- a/qa/L0_pytorch_lint/test.sh +++ b/qa/L0_pytorch_lint/test.sh @@ -15,6 +15,8 @@ then echo "Checking C++ files" python3 -m cpplint --recursive --exclude=transformer_engine/common/include --exclude=transformer_engine/build_tools/build transformer_engine/common python3 -m cpplint --recursive transformer_engine/pytorch + echo "Checking TE<->PyTorch binary boundary" + python3 qa/L0_pytorch_lint/check_torch_boundary.py "$TE_PATH" fi if [ -z "${CPP_ONLY}" ] then diff --git a/qa/L0_pytorch_lint/torch_boundary_allowlist.txt b/qa/L0_pytorch_lint/torch_boundary_allowlist.txt new file mode 100644 index 0000000000..3819fa313f --- /dev/null +++ b/qa/L0_pytorch_lint/torch_boundary_allowlist.txt @@ -0,0 +1,11 @@ +# Files in transformer_engine/pytorch/csrc that still talk to the PyTorch ABI +# (at::/c10::/torch::) directly instead of going through torch_backend.h. +# +# This list only shrinks: when a file is migrated to the torch_backend.h +# facade, remove it here. The guard (check_torch_boundary.py) fails if a listed +# file no longer touches the torch ABI (stale) or if an unlisted file starts to. +# +# Paths are relative to transformer_engine/pytorch/csrc/. +# +# STATUS: empty -- the entire pytorch/csrc extension now routes every torch/ATen/ +# c10 access through torch_backend.h. Nothing is grandfathered. Keep it that way. diff --git a/transformer_engine/pytorch/csrc/common.cpp b/transformer_engine/pytorch/csrc/common.cpp index d85dcda159..08e07242a0 100644 --- a/transformer_engine/pytorch/csrc/common.cpp +++ b/transformer_engine/pytorch/csrc/common.cpp @@ -40,7 +40,7 @@ std::array get_2d_dims(NVTEShape shape, bool transpose) { } } -std::vector getTensorShape(const at::Tensor& t) { +std::vector getTensorShape(const Tensor& t) { std::vector shape; for (auto s : t.sizes()) { shape.push_back(s); @@ -48,7 +48,7 @@ std::vector getTensorShape(const at::Tensor& t) { return shape; } -NVTEShape convertTorchShape(const c10::IntArrayRef torch_shape) { +NVTEShape convertTorchShape(const IntArrayRef torch_shape) { NVTEShape ret; ret.ndim = torch_shape.size(); constexpr int max_dimensions = sizeof(ret.data) / sizeof(size_t); @@ -132,7 +132,7 @@ TensorWrapper makeTransformerEngineTensor(py::handle tensor, py::handle quantize "Unexpected quantization params type."); // Regular pyTorch tensor - at::Tensor torch_tensor = tensor.cast(); + Tensor torch_tensor = tensor.cast(); // #TODO (pgadzinski) - needed in attention for non-contiguous tensors. //if (!torch_tensor.is_contiguous()) { @@ -156,7 +156,7 @@ transformer_engine::TensorWrapper makeTransformerEngineTensor( return transformer_engine::TensorWrapper(data_ptr, shape, type); } -transformer_engine::TensorWrapper makeTransformerEngineTensor(at::Tensor tensor) { +transformer_engine::TensorWrapper makeTransformerEngineTensor(Tensor tensor) { transformer_engine::DType dtype = GetTransformerEngineDType(tensor.scalar_type()); std::vector shape; for (auto s : tensor.sizes()) { @@ -167,7 +167,7 @@ transformer_engine::TensorWrapper makeTransformerEngineTensor(at::Tensor tensor) std::tuple, std::vector>, std::vector, size_t, size_t> -makeTransformerEngineTensorList(std::vector> at_tensor_lists) { +makeTransformerEngineTensorList(std::vector> at_tensor_lists) { size_t num_lists = at_tensor_lists.size(); NVTE_CHECK(num_lists > 0, "List of tensors is empty."); @@ -238,18 +238,18 @@ transformer_engine::TensorWrapper makeTransformerEngineTensor( return ret; } -transformer_engine::TensorWrapper makeTransformerEngineTensor(at::Tensor tensor, at::Tensor amax, - const at::Tensor scale, - at::Tensor scale_inv, +transformer_engine::TensorWrapper makeTransformerEngineTensor(Tensor tensor, Tensor amax, + const Tensor scale, + Tensor scale_inv, NVTEScalingMode scaling_mode) { transformer_engine::DType dtype = GetTransformerEngineDType(tensor.scalar_type()); auto tensor_shape = getTensorShape(tensor); auto scale_inv_shape = getTensorShape(scale_inv); - NVTE_CHECK(amax.scalar_type() == at::kFloat); - NVTE_CHECK(scale.scalar_type() == at::kFloat); - NVTE_CHECK(scale_inv.scalar_type() == at::kFloat); + NVTE_CHECK(GetTransformerEngineDType(amax.scalar_type()) == DType::kFloat32); + NVTE_CHECK(GetTransformerEngineDType(scale.scalar_type()) == DType::kFloat32); + NVTE_CHECK(GetTransformerEngineDType(scale_inv.scalar_type()) == DType::kFloat32); return makeTransformerEngineTensor(tensor.data_ptr(), tensor_shape, dtype, amax.data_ptr(), scale.data_ptr(), scale_inv.data_ptr(), scale_inv_shape, @@ -286,45 +286,37 @@ std::vector nvte_shape_to_vector(const NVTEShape& nvte_shape) { return shape; } -at::Tensor allocateSpace(const std::vector& shape, const transformer_engine::DType type, +Tensor allocateSpace(const std::vector& shape, const transformer_engine::DType type, bool init_to_zeros) { std::vector shape_int64(shape.begin(), shape.end()); - c10::IntArrayRef ar_shape(shape_int64); - if (init_to_zeros) { - return at::zeros(ar_shape, at::CUDA(GetATenDType(type))); - } else { - return at::empty(ar_shape, at::CUDA(GetATenDType(type))); - } + return new_cuda_tensor(shape_int64, GetATenDType(type), init_to_zeros); } -at::Tensor allocateSpace(const NVTEShape& shape, const transformer_engine::DType type, +Tensor allocateSpace(const NVTEShape& shape, const transformer_engine::DType type, bool init_to_zeros) { auto size = shape.ndim; - if (size == 2 && init_to_zeros) { - return at::zeros({static_cast(shape.data[0]), static_cast(shape.data[1])}, - at::CUDA(GetATenDType(type))); - } else if (size == 2) { - return at::empty({static_cast(shape.data[0]), static_cast(shape.data[1])}, - at::CUDA(GetATenDType(type))); - } else if (size == 1 && init_to_zeros) { - return at::zeros({static_cast(shape.data[0])}, at::CUDA(GetATenDType(type))); + if (size == 2) { + return new_cuda_tensor( + {static_cast(shape.data[0]), static_cast(shape.data[1])}, + GetATenDType(type), init_to_zeros); } else if (size == 1) { - return at::empty({static_cast(shape.data[0])}, at::CUDA(GetATenDType(type))); + return new_cuda_tensor({static_cast(shape.data[0])}, GetATenDType(type), + init_to_zeros); } NVTE_ERROR("Unsupported tensor allocation: ndim=", size, ", init_to_zeros=", init_to_zeros, ". Only 1D and 2D tensors are supported."); } -at::Tensor allocateTorchTensor(int M, int N, transformer_engine::DType dtype) { - return at::empty({static_cast(M), static_cast(N)}, - at::CUDA(GetATenDType(dtype))); +Tensor allocateTorchTensor(int M, int N, transformer_engine::DType dtype) { + return new_cuda_tensor({static_cast(M), static_cast(N)}, GetATenDType(dtype), + /*zero_init=*/false); } -at::Tensor allocateTorchTensor(int M, transformer_engine::DType dtype) { - return at::empty({static_cast(M)}, at::CUDA(GetATenDType(dtype))); +Tensor allocateTorchTensor(int M, transformer_engine::DType dtype) { + return new_cuda_tensor({static_cast(M)}, GetATenDType(dtype), /*zero_init=*/false); } -void* getDataPtr(at::Tensor tensor, int offset) { +void* getDataPtr(Tensor tensor, int offset) { void* dptr = nullptr; if (tensor.numel() > 0) { dptr = tensor.data_ptr(); @@ -348,17 +340,17 @@ size_t roundup(size_t value, size_t multiple) { size_t ceildiv(size_t numer, size_t denom) { return (numer + denom - 1) / denom; } -void philox_unpack(at::PhiloxCudaState arg, int64_t* rng_state_ptr) { +void philox_unpack(PhiloxCudaState arg, int64_t* rng_state_ptr) { NVTE_SCOPED_GIL_RELEASE({ nvte_extract_seed_and_offset(rng_state_ptr, arg.captured_, arg.seed_.ptr, arg.seed_.val, arg.offset_.ptr, arg.offset_.val, arg.offset_intragraph_, - at::cuda::getCurrentCUDAStream()); + getCurrentCUDAStream()); }); } // extract PhiloxCudaState from CUDA random number generator -at::PhiloxCudaState init_philox_state(at::CUDAGeneratorImpl* gen, size_t elts_per_thread) { - at::PhiloxCudaState philox_args; +PhiloxCudaState init_philox_state(CUDAGeneratorImpl* gen, size_t elts_per_thread) { + PhiloxCudaState philox_args; std::lock_guard lock(gen->mutex_); philox_args = gen->philox_cuda_state(elts_per_thread); return philox_args; diff --git a/transformer_engine/pytorch/csrc/common.h b/transformer_engine/pytorch/csrc/common.h index 779b145dd9..6989e86925 100644 --- a/transformer_engine/pytorch/csrc/common.h +++ b/transformer_engine/pytorch/csrc/common.h @@ -7,22 +7,11 @@ #ifndef TRANSFORMER_ENGINE_PYTORCH_CSRC_COMMON_H_ #define TRANSFORMER_ENGINE_PYTORCH_CSRC_COMMON_H_ -#include -#include -#include -#include -#include -#include -#include -#include -#include #include #include #include #include #include -#include -#include #include #include #include @@ -44,31 +33,31 @@ #include #include -#include #include #include #include #include #include -#include #include -#include "c10/util/ArrayRef.h" #include "common/util/logging.h" #include "extensions/pybind_dtype_caster.h" +// Single binary boundary with PyTorch: all torch/ATen/c10 types and ops used +// below come from here as aliases (Tensor, Device, ScalarType, ...). +#include "torch_backend.h" namespace transformer_engine::pytorch { // in python we have: dist_group_type = torch.distributed.ProcessGroup -using dist_group_type = c10d::ProcessGroup; +using dist_group_type = ProcessGroup; // Each tensor here is shape (N, ) holding all scaling // data for a single FP8 block, e.g. LayerNormLinear class FP8TensorMeta { public: - at::Tensor scale; - at::Tensor scale_inv; - at::Tensor amax_history; + Tensor scale; + Tensor scale_inv; + Tensor amax_history; }; // Used as named indices on the `scale`, `scale_inv`, @@ -105,7 +94,7 @@ class Quantizer { /*! @brief Construct a tensor with uninitialized data */ virtual std::pair create_tensor( const std::vector& shape, DType dtype, - std::optional device = std::nullopt, bool pin_memory = false) const = 0; + std::optional device = std::nullopt, bool pin_memory = false) const = 0; /*! @brief Construct a grouped tensor with uninitialized data * @@ -118,8 +107,8 @@ class Quantizer { */ virtual std::pair create_grouped_tensor( size_t num_tensors, const std::vector& logical_shape, DType dtype, - py::object quantizer, const std::optional& first_dims, - const std::optional& tensor_offsets, size_t logical_first_dim, + py::object quantizer, const std::optional& first_dims, + const std::optional& tensor_offsets, size_t logical_first_dim, size_t logical_last_dim) const = 0; /*! @brief Convert a PyTorch tensor into a Transformer Engine C++ tensor @@ -157,17 +146,17 @@ class NoneQuantizer : public Quantizer { std::pair create_tensor( const std::vector& shape, DType dtype, - std::optional device = std::nullopt, bool pin_memory = false) const override; + std::optional device = std::nullopt, bool pin_memory = false) const override; std::pair create_grouped_tensor( size_t num_tensors, const std::vector& logical_shape, DType dtype, - py::object quantizer, const std::optional& first_dims, - const std::optional& tensor_offsets, size_t logical_first_dim, + py::object quantizer, const std::optional& first_dims, + const std::optional& tensor_offsets, size_t logical_first_dim, size_t logical_last_dim) const override; /*! @brief Construct a tensor with pre-initialized data */ std::pair create_tensor(const std::vector& shape, DType dtype, - at::Tensor data) const; + Tensor data) const; std::pair convert_and_update_tensor(py::object tensor) const override; @@ -177,9 +166,9 @@ class NoneQuantizer : public Quantizer { class Float8Quantizer : public Quantizer { public: - at::Tensor scale; - at::Tensor scale_inv; - at::Tensor amax; + Tensor scale; + Tensor scale_inv; + Tensor amax; explicit Float8Quantizer(const py::handle& quantizer); @@ -189,19 +178,19 @@ class Float8Quantizer : public Quantizer { std::pair create_tensor( const std::vector& shape, DType dtype, - std::optional device = std::nullopt, bool pin_memory = false) const override; + std::optional device = std::nullopt, bool pin_memory = false) const override; std::pair create_grouped_tensor( size_t num_tensors, const std::vector& logical_shape, DType dtype, - py::object quantizer, const std::optional& first_dims, - const std::optional& tensor_offsets, size_t logical_first_dim, + py::object quantizer, const std::optional& first_dims, + const std::optional& tensor_offsets, size_t logical_first_dim, size_t logical_last_dim) const override; /*! @brief Construct a tensor with pre-initialized data */ std::pair create_tensor( - const std::vector& shape, DType dtype, std::optional data, - std::optional transpose, std::optional scale_inv, - std::optional device = std::nullopt, bool pin_memory = false) const; + const std::vector& shape, DType dtype, std::optional data, + std::optional transpose, std::optional scale_inv, + std::optional device = std::nullopt, bool pin_memory = false) const; std::pair convert_and_update_tensor(py::object shape) const override; @@ -213,7 +202,7 @@ class Float8CurrentScalingQuantizer : public Quantizer { public: DType dtype; bool with_amax_reduction; - c10::intrusive_ptr amax_reduction_group; + IntrusivePtr amax_reduction_group; bool force_pow_2_scales = false; float amax_epsilon = 0.0; @@ -225,12 +214,12 @@ class Float8CurrentScalingQuantizer : public Quantizer { std::pair create_tensor( const std::vector& shape, DType dtype, - std::optional device = std::nullopt, bool pin_memory = false) const override; + std::optional device = std::nullopt, bool pin_memory = false) const override; std::pair create_grouped_tensor( size_t num_tensors, const std::vector& logical_shape, DType dtype, - py::object quantizer, const std::optional& first_dims, - const std::optional& tensor_offsets, size_t logical_first_dim, + py::object quantizer, const std::optional& first_dims, + const std::optional& tensor_offsets, size_t logical_first_dim, size_t logical_last_dim) const override; /*! @brief Construct an unquantized tensor with a freshly allocated amax buffer. @@ -239,8 +228,8 @@ class Float8CurrentScalingQuantizer : public Quantizer { * amax to be initialized to zero. The amax tensor is returned as * the third element to keep it alive in the caller's scope. */ - std::tuple create_unquantized_tensor_with_amax( - const std::vector& shape, DType dtype, std::optional data = std::nullopt); + std::tuple create_unquantized_tensor_with_amax( + const std::vector& shape, DType dtype, std::optional data = std::nullopt); std::pair convert_and_update_tensor(py::object shape) const override; @@ -253,13 +242,13 @@ class Float8CurrentScalingQuantizer : public Quantizer { * amax. The amax may still be reduced across the amax reduction * group. */ - void quantize_with_amax(TensorWrapper& input, TensorWrapper& out, at::Tensor amax, + void quantize_with_amax(TensorWrapper& input, TensorWrapper& out, Tensor amax, const std::optional& noop_flag = std::nullopt); private: void quantize_impl(const TensorWrapper& input, TensorWrapper& out, const std::optional& noop_flag, bool compute_amax, - at::Tensor amax_buf, at::Tensor scale_buf); + Tensor amax_buf, Tensor scale_buf); }; class Float8BlockQuantizer : public Quantizer { @@ -289,12 +278,12 @@ class Float8BlockQuantizer : public Quantizer { // and optionally columnwise usage. std::pair create_tensor( const std::vector& shape, DType dtype, - std::optional device = std::nullopt, bool pin_memory = false) const override; + std::optional device = std::nullopt, bool pin_memory = false) const override; std::pair create_grouped_tensor( size_t num_tensors, const std::vector& logical_shape, DType dtype, - py::object quantizer, const std::optional& first_dims, - const std::optional& tensor_offsets, size_t logical_first_dim, + py::object quantizer, const std::optional& first_dims, + const std::optional& tensor_offsets, size_t logical_first_dim, size_t logical_last_dim) const override; std::pair convert_and_update_tensor(py::object shape) const override; @@ -315,12 +304,12 @@ class MXFP8Quantizer : public Quantizer { std::pair create_tensor( const std::vector& shape, DType dtype, - std::optional device = std::nullopt, bool pin_memory = false) const override; + std::optional device = std::nullopt, bool pin_memory = false) const override; std::pair create_grouped_tensor( size_t num_tensors, const std::vector& logical_shape, DType dtype, - py::object quantizer, const std::optional& first_dims, - const std::optional& tensor_offsets, size_t logical_first_dim, + py::object quantizer, const std::optional& first_dims, + const std::optional& tensor_offsets, size_t logical_first_dim, size_t logical_last_dim) const override; std::pair convert_and_update_tensor(py::object shape) const override; @@ -335,7 +324,7 @@ class NVFP4Quantizer : public Quantizer { public: // amax reduction for low precision FP4 AG bool with_amax_reduction; - c10::intrusive_ptr amax_reduction_group; + IntrusivePtr amax_reduction_group; // random hadamard transform bool with_rht; bool with_post_rht_amax; @@ -350,7 +339,7 @@ class NVFP4Quantizer : public Quantizer { bool row_scaled_nvfp4; int rht_matrix_random_sign_mask_t; - at::Tensor rht_matrix; + Tensor rht_matrix; explicit NVFP4Quantizer(const py::handle& quantizer); @@ -360,12 +349,12 @@ class NVFP4Quantizer : public Quantizer { std::pair create_tensor( const std::vector& shape, DType dtype, - std::optional device = std::nullopt, bool pin_memory = false) const override; + std::optional device = std::nullopt, bool pin_memory = false) const override; std::pair create_grouped_tensor( size_t num_tensors, const std::vector& logical_shape, DType dtype, - py::object quantizer, const std::optional& first_dims, - const std::optional& tensor_offsets, size_t logical_first_dim, + py::object quantizer, const std::optional& first_dims, + const std::optional& tensor_offsets, size_t logical_first_dim, size_t logical_last_dim) const override; /*! @brief Construct an unquantized tensor that shares NVFP4 tensor's amax pointer @@ -409,7 +398,7 @@ class NVFP4Quantizer : public Quantizer { std::unique_ptr convert_quantizer(py::handle quantizer); -std::vector getTensorShape(const at::Tensor& t); +std::vector getTensorShape(const Tensor& t); transformer_engine::DType getTransformerEngineFP8Type(bool e4m3_if_hybrid, const std::string& fp8_recipe); @@ -446,59 +435,8 @@ inline size_t typeToNumBits(transformer_engine::DType t) { } } -inline at::ScalarType GetATenDType(transformer_engine::DType t) { - switch (t) { - case transformer_engine::DType::kInt16: - return torch::kInt16; - case transformer_engine::DType::kInt32: - return torch::kInt32; - case transformer_engine::DType::kInt64: - return torch::kInt64; - case transformer_engine::DType::kFloat32: - return at::kFloat; - case transformer_engine::DType::kFloat16: - return at::kHalf; - case transformer_engine::DType::kBFloat16: - return at::kBFloat16; - case transformer_engine::DType::kByte: - return at::kByte; - case transformer_engine::DType::kFloat8E4M3: - return at::kFloat8_e4m3fn; - case transformer_engine::DType::kFloat8E5M2: - return at::kFloat8_e5m2; - case transformer_engine::DType::kFloat8E8M0: - return at::kByte; // e8m0 dtype requires PyTorch 2.7.0+ - default: - NVTE_ERROR("Invalid type (", static_cast(t), ")."); - } -} - -inline transformer_engine::DType GetTransformerEngineDType(at::ScalarType t) { - switch (t) { - case at::kFloat8_e4m3fn: - return transformer_engine::DType::kFloat8E4M3; - case at::kFloat8_e5m2: - return transformer_engine::DType::kFloat8E5M2; - case at::kHalf: - return transformer_engine::DType::kFloat16; - case at::kFloat: - return transformer_engine::DType::kFloat32; - case at::kBFloat16: - return transformer_engine::DType::kBFloat16; - case at::kBool: - return transformer_engine::DType::kByte; - case torch::kByte: - return transformer_engine::DType::kByte; - case torch::kInt16: - return transformer_engine::DType::kInt16; - case torch::kInt32: - return transformer_engine::DType::kInt32; - case torch::kInt64: - return transformer_engine::DType::kInt64; - default: - NVTE_ERROR("Invalid type (", static_cast(t), ")."); - } -} +// GetATenDType and GetTransformerEngineDType(ScalarType) are the TE<->torch +// dtype boundary and live in torch_backend.h / torch_backend.cpp. inline transformer_engine::DType GetTransformerEngineDType(int DType_value) { return static_cast(DType_value); @@ -525,16 +463,16 @@ transformer_engine::TensorWrapper makeTransformerEngineTensor(void* data_ptr, const NVTEShape& shape, const transformer_engine::DType type); -transformer_engine::TensorWrapper makeTransformerEngineTensor(at::Tensor tensor); +transformer_engine::TensorWrapper makeTransformerEngineTensor(Tensor tensor); std::tuple, std::vector>, std::vector, size_t, size_t> -makeTransformerEngineTensorList(std::vector> at_tensor_lists); +makeTransformerEngineTensorList(std::vector> at_tensor_lists); TensorWrapper makeTransformerEngineTensor(py::handle tensor, py::handle quantizer); transformer_engine::TensorWrapper makeTransformerEngineTensor( - at::Tensor tensor, at::Tensor amax, const at::Tensor scale, at::Tensor scale_inv, + Tensor tensor, Tensor amax, const Tensor scale, Tensor scale_inv, NVTEScalingMode scaling_mode = NVTE_DELAYED_TENSOR_SCALING); template @@ -544,17 +482,17 @@ size_t product(const NVTEShape& shape, size_t begin, size_t end); std::vector nvte_shape_to_vector(const NVTEShape& nvte_shape); -at::Tensor allocateSpace(const std::vector& shape, const transformer_engine::DType type, +Tensor allocateSpace(const std::vector& shape, const transformer_engine::DType type, bool init_to_zeros); -at::Tensor allocateSpace(const NVTEShape& shape, const transformer_engine::DType type, +Tensor allocateSpace(const NVTEShape& shape, const transformer_engine::DType type, bool init_to_zeros = false); -at::Tensor allocateTorchTensor(int M, int N, transformer_engine::DType dtype); +Tensor allocateTorchTensor(int M, int N, transformer_engine::DType dtype); -at::Tensor allocateTorchTensor(int M, transformer_engine::DType dtype); +Tensor allocateTorchTensor(int M, transformer_engine::DType dtype); -void* getDataPtr(at::Tensor tensor, int offset = 0); +void* getDataPtr(Tensor tensor, int offset = 0); std::vector convertShape(const NVTEShape& shape); @@ -562,7 +500,7 @@ size_t roundup(size_t value, size_t multiple); size_t ceildiv(size_t numer, size_t denom); -NVTEShape convertTorchShape(const c10::IntArrayRef torch_shape); +NVTEShape convertTorchShape(const IntArrayRef torch_shape); std::vector convert_shape_back_from_fp4(const std::vector& shape, bool transpose); @@ -582,10 +520,10 @@ inline std::array get_2d_dims(const std::vector& shape, bool trans } // unpack the PhiloxCudaState into CUDA tensor -void philox_unpack(at::PhiloxCudaState arg, int64_t* rng_state_ptr); +void philox_unpack(PhiloxCudaState arg, int64_t* rng_state_ptr); // extract PhiloxCudaState from CUDA random number generator -at::PhiloxCudaState init_philox_state(at::CUDAGeneratorImpl* gen, size_t elts_per_thread); +PhiloxCudaState init_philox_state(CUDAGeneratorImpl* gen, size_t elts_per_thread); } // namespace transformer_engine::pytorch @@ -604,9 +542,10 @@ string to_string(const vector& vec) { return ret; } -// Torch shape -> string +// Torch shape -> string. Fully qualified because the ArrayRef alias lives in +// transformer_engine::pytorch and this overload is in namespace std. template -string to_string(const c10::ArrayRef& vec) { +string to_string(const transformer_engine::pytorch::ArrayRef& vec) { string ret = "["; for (const auto& val : vec) { ret += to_string(val) + ","; diff --git a/transformer_engine/pytorch/csrc/extensions.h b/transformer_engine/pytorch/csrc/extensions.h index 2b4f899e1d..c9e05595e1 100644 --- a/transformer_engine/pytorch/csrc/extensions.h +++ b/transformer_engine/pytorch/csrc/extensions.h @@ -30,51 +30,51 @@ namespace transformer_engine::pytorch { * Router fusion **************************************************************************************************/ -std::tuple fused_topk_with_score_function_fwd( - at::Tensor logits, int topk, bool use_pre_softmax, std::optional num_groups, +std::tuple fused_topk_with_score_function_fwd( + Tensor logits, int topk, bool use_pre_softmax, std::optional num_groups, std::optional group_topk, std::optional scaling_factor, std::string score_function, - std::optional expert_bias, + std::optional expert_bias, int routing_map_format = static_cast(NVTE_ROUTING_MAP_FORMAT_BYTEMAP)); void fused_topk_with_score_function_bwd( - at::Tensor routing_map, at::Tensor intermediate_output, at::Tensor grad_probs, - at::Tensor grad_logits, int topk, bool use_pre_softmax, std::optional scaling_factor, + Tensor routing_map, Tensor intermediate_output, Tensor grad_probs, + Tensor grad_logits, int topk, bool use_pre_softmax, std::optional scaling_factor, std::string score_function, int routing_map_format = static_cast(NVTE_ROUTING_MAP_FORMAT_BYTEMAP)); -std::tuple fused_score_for_moe_aux_loss_fwd( - at::Tensor logits, int topk, std::string score_function, +std::tuple fused_score_for_moe_aux_loss_fwd( + Tensor logits, int topk, std::string score_function, int routing_map_format = static_cast(NVTE_ROUTING_MAP_FORMAT_BYTEMAP)); -void fused_score_for_moe_aux_loss_bwd(at::Tensor intermediate_output, at::Tensor grad_scores, - at::Tensor grad_logits, int topk, std::string score_function); +void fused_score_for_moe_aux_loss_bwd(Tensor intermediate_output, Tensor grad_scores, + Tensor grad_logits, int topk, std::string score_function); -std::tuple fused_moe_aux_loss_fwd(at::Tensor probs, - at::Tensor tokens_per_expert, +std::tuple fused_moe_aux_loss_fwd(Tensor probs, + Tensor tokens_per_expert, int total_num_tokens, int num_experts, int num_rows, int num_cols, int topk, float coeff); -at::Tensor fused_moe_aux_loss_bwd(at::Tensor Const_buf, at::Tensor tokens_per_expert, int num_rows, - int num_cols, at::Tensor grad_aux_loss); +Tensor fused_moe_aux_loss_bwd(Tensor Const_buf, Tensor tokens_per_expert, int num_rows, + int num_cols, Tensor grad_aux_loss); /*************************************************************************************************** * Permutation **************************************************************************************************/ -std::tuple> moe_permute_fwd( - at::Tensor input, const DType dtype, at::Tensor indices, int64_t num_out_tokens, - std::vector workspace, int64_t max_expanded_token_num); +std::tuple> moe_permute_fwd( + Tensor input, const DType dtype, Tensor indices, int64_t num_out_tokens, + std::vector workspace, int64_t max_expanded_token_num); -at::Tensor moe_permute_bwd(at::Tensor input, const DType dtype, at::Tensor row_id_map, - at::Tensor prob, int64_t num_tokens, int64_t topK); +Tensor moe_permute_bwd(Tensor input, const DType dtype, Tensor row_id_map, + Tensor prob, int64_t num_tokens, int64_t topK); -at::Tensor moe_unpermute_fwd(at::Tensor input, const DType dtype, at::Tensor row_id_map, - at::Tensor prob, int64_t num_tokens, int64_t topK); +Tensor moe_unpermute_fwd(Tensor input, const DType dtype, Tensor row_id_map, + Tensor prob, int64_t num_tokens, int64_t topK); -std::tuple moe_unpermute_bwd(at::Tensor input_bwd, at::Tensor input_fwd, - const DType dtype, at::Tensor row_id_map, - at::Tensor prob); +std::tuple moe_unpermute_bwd(Tensor input_bwd, Tensor input_fwd, + const DType dtype, Tensor row_id_map, + Tensor prob); /*************************************************************************************************** * Attention @@ -92,13 +92,13 @@ std::vector fused_attn_fwd( bool set_zero, NVTE_QKV_Layout qkv_layout, NVTE_QKV_Format o_format, NVTE_QKV_Format qkv_scale_inv_format, NVTE_Bias_Type bias_type, NVTE_Mask_Type attn_mask_type, NVTE_Softmax_Type softmax_type, const std::vector window_size, - bool bottom_right_diagonal, const at::Tensor cu_seqlens_q, const at::Tensor cu_seqlens_kv, - const py::handle Q, const py::handle K, const py::handle V, const at::ScalarType fake_dtype, - const std::optional cu_seqlens_q_padded, - const std::optional cu_seqlens_kv_padded, - const std::optional page_table_k, const std::optional page_table_v, - py::handle s_quantizer, py::handle o_quantizer, const std::optional Bias, - const std::optional SoftmaxOffset, const std::optional rng_gen, + bool bottom_right_diagonal, const Tensor cu_seqlens_q, const Tensor cu_seqlens_kv, + const py::handle Q, const py::handle K, const py::handle V, const ScalarType fake_dtype, + const std::optional cu_seqlens_q_padded, + const std::optional cu_seqlens_kv_padded, + const std::optional page_table_k, const std::optional page_table_v, + py::handle s_quantizer, py::handle o_quantizer, const std::optional Bias, + const std::optional SoftmaxOffset, const std::optional rng_gen, size_t rng_elts_per_thread, bool return_max_logit, bool cuda_graph); std::vector fused_attn_bwd( @@ -107,28 +107,28 @@ std::vector fused_attn_bwd( NVTE_QKV_Layout dqkv_layout, NVTE_QKV_Format qkv_scale_inv_format, NVTE_QKV_Format do_scale_inv_format, NVTE_Bias_Type bias_type, NVTE_Mask_Type attn_mask_type, NVTE_Softmax_Type softmax_type, const std::vector window_size, - bool bottom_right_diagonal, bool deterministic, const at::Tensor cu_seqlens_q, - const at::Tensor cu_seqlens_kv, const py::handle Q, const py::handle K, const py::handle V, - const py::handle O, const py::handle dO, const at::ScalarType fake_dtype, - const std::vector Aux_CTX_Tensors, - const std::optional cu_seqlens_q_padded, - const std::optional cu_seqlens_kv_padded, py::handle s_quantizer, + bool bottom_right_diagonal, bool deterministic, const Tensor cu_seqlens_q, + const Tensor cu_seqlens_kv, const py::handle Q, const py::handle K, const py::handle V, + const py::handle O, const py::handle dO, const ScalarType fake_dtype, + const std::vector Aux_CTX_Tensors, + const std::optional cu_seqlens_q_padded, + const std::optional cu_seqlens_kv_padded, py::handle s_quantizer, py::handle dp_quantizer, py::handle dqkv_quantizer, bool cuda_graph); -at::Tensor fa_prepare_fwd(at::Tensor qkvi); -at::Tensor fa_prepare_bwd(at::Tensor q, at::Tensor k, at::Tensor v); +Tensor fa_prepare_fwd(Tensor qkvi); +Tensor fa_prepare_bwd(Tensor q, Tensor k, Tensor v); -std::vector> multi_tensor_transpose_to_bhsd( - std::vector> inputs, const std::string &original_format, - std::vector> outputs = {}); +std::vector> multi_tensor_transpose_to_bhsd( + std::vector> inputs, const std::string &original_format, + std::vector> outputs = {}); -std::vector multi_tensor_pad_last_dim(std::vector inputs, +std::vector multi_tensor_pad_last_dim(std::vector inputs, int64_t alignment); -at::Tensor convert_thd_to_bshd(at::Tensor tensor, at::Tensor cu_seqlens, int b, int max_seq_len); -at::Tensor convert_bshd_to_thd(at::Tensor tensor, at::Tensor cu_seqlens, int t); -void copy_to_kv_cache(at::Tensor new_k, at::Tensor new_v, at::Tensor k_cache, at::Tensor v_cache, - at::Tensor page_table, at::Tensor cu_new_lens, at::Tensor cu_cached_lens, +Tensor convert_thd_to_bshd(Tensor tensor, Tensor cu_seqlens, int b, int max_seq_len); +Tensor convert_bshd_to_thd(Tensor tensor, Tensor cu_seqlens, int t); +void copy_to_kv_cache(Tensor new_k, Tensor new_v, Tensor k_cache, Tensor v_cache, + Tensor page_table, Tensor cu_new_lens, Tensor cu_cached_lens, NVTE_QKV_Format kv_format, int b, int max_ctx_len, int max_seq_len, int max_pages_per_seq, bool is_non_paged); @@ -136,161 +136,161 @@ void copy_to_kv_cache(at::Tensor new_k, at::Tensor new_v, at::Tensor k_cache, at * GEMM **************************************************************************************************/ -using MaybeTensor = std::optional; +using MaybeTensor = std::optional; std::vector gemm(py::handle A, bool transa, py::handle B, bool transb, py::object D, py::handle quantizer, std::optional out_dtype, MaybeTensor bias, DType bias_type, bool gelu, MaybeTensor gelu_in, bool grad, - at::Tensor workspace, size_t workspaceSize, bool accumulate, + Tensor workspace, size_t workspaceSize, bool accumulate, bool use_split_accumulator, CommOverlapCore *comm_overlap = nullptr, std::optional comm_type = std::nullopt, MaybeTensor extra_output = std::nullopt, bool bulk_overlap = false, float alpha = 1.0f, std::optional beta = std::nullopt); -void te_atomic_gemm(at::Tensor A, at::Tensor A_scale_inverse, DType A_type, - std::vector A_scaling_mode, bool transa, at::Tensor B, - at::Tensor B_scale_inverse, DType B_type, std::vector B_scaling_mode, - bool transb, at::Tensor D, at::Tensor D_scale, DType D_type, at::Tensor D_amax, - at::Tensor bias, DType bias_type, at::Tensor pre_gelu_out, bool grad, - at::Tensor workspace, size_t workspaceSize, bool accumulate, +void te_atomic_gemm(Tensor A, Tensor A_scale_inverse, DType A_type, + std::vector A_scaling_mode, bool transa, Tensor B, + Tensor B_scale_inverse, DType B_type, std::vector B_scaling_mode, + bool transb, Tensor D, Tensor D_scale, DType D_type, Tensor D_amax, + Tensor bias, DType bias_type, Tensor pre_gelu_out, bool grad, + Tensor workspace, size_t workspaceSize, bool accumulate, bool use_split_accumulator, int math_sm_count, int m_split, int n_split, - bool gemm_producer, at::Tensor counter); + bool gemm_producer, Tensor counter); -std::optional> te_general_grouped_gemm( +std::optional> te_general_grouped_gemm( std::vector A, bool transa, std::vector B, bool transb, - std::optional> D, DType D_type, std::vector m_splits, - std::vector bias, DType bias_type, bool single_output, - std::vector pre_gelu_out, bool grad, std::vector workspace, + std::optional> D, DType D_type, std::vector m_splits, + std::vector bias, DType bias_type, bool single_output, + std::vector pre_gelu_out, bool grad, std::vector workspace, size_t workspaceSize, bool accumulate, bool use_split_accumulator, int math_sm_count); py::object te_general_grouped_gemm_for_grouped_tensor( py::handle A, bool transa, py::handle B, bool transb, py::handle D, py::object bias, - std::optional bias_scale, at::Tensor alpha, at::Tensor beta, - at::Tensor workspace_setup, at::Tensor workspace_cublas, bool use_split_accumulator, + std::optional bias_scale, Tensor alpha, Tensor beta, + Tensor workspace_setup, Tensor workspace_cublas, bool use_split_accumulator, int math_sm_count); py::object te_general_grouped_gemm_for_discrete_in(py::handle A, bool transa, py::handle B, bool transb, py::handle D, py::object bias, - std::optional bias_scale, - at::Tensor alpha, at::Tensor beta, - at::Tensor workspace_setup, - at::Tensor workspace_cublas, + std::optional bias_scale, + Tensor alpha, Tensor beta, + Tensor workspace_setup, + Tensor workspace_cublas, bool use_split_accumulator, int math_sm_count); py::object te_general_grouped_gemm_for_discrete_out(py::handle A, bool transa, py::handle B, bool transb, py::handle D, py::object bias, - std::optional bias_scale, - at::Tensor alpha, at::Tensor beta, - at::Tensor workspace_setup, - at::Tensor workspace_cublas, + std::optional bias_scale, + Tensor alpha, Tensor beta, + Tensor workspace_setup, + Tensor workspace_cublas, bool use_split_accumulator, int math_sm_count); /*************************************************************************************************** * Transpose **************************************************************************************************/ -at::Tensor fp8_transpose(at::Tensor input, DType otype, - std::optional output = std::nullopt); +Tensor fp8_transpose(Tensor input, DType otype, + std::optional output = std::nullopt); -at::Tensor nvfp4_data_transpose(at::Tensor input, std::optional output = std::nullopt); +Tensor nvfp4_data_transpose(Tensor input, std::optional output = std::nullopt); -void nvfp4_2d_scale_transpose(at::Tensor input, at::Tensor output, int64_t M_tiles, +void nvfp4_2d_scale_transpose(Tensor input, Tensor output, int64_t M_tiles, int64_t K_tiles); -void nvfp4_2d_multi_tensor_transpose(std::vector rowwise_data_list, - std::vector columnwise_data_list, - std::vector rowwise_scale_inv_list, - std::vector columnwise_scale_inv_list, +void nvfp4_2d_multi_tensor_transpose(std::vector rowwise_data_list, + std::vector columnwise_data_list, + std::vector rowwise_scale_inv_list, + std::vector columnwise_scale_inv_list, std::vector M_list, std::vector K_list); void nvfp4_multi_tensor_compute_partial_amax( - std::vector master_weight_list, std::vector partial_amax_list, - std::vector global_amax_list, std::vector h_list, + std::vector master_weight_list, std::vector partial_amax_list, + std::vector global_amax_list, std::vector h_list, std::vector w_list, std::vector start_offset_list, int64_t block_len); -void nvfp4_expand_scale_to_fp8(at::Tensor input, at::Tensor output, int64_t tile_rows, +void nvfp4_expand_scale_to_fp8(Tensor input, Tensor output, int64_t tile_rows, int64_t tile_cols, int64_t rows_padded, int64_t block_len); -void nvfp4_compute_per_block_scale(at::Tensor block_amax, at::Tensor scale, at::Tensor global_amax); +void nvfp4_compute_per_block_scale(Tensor block_amax, Tensor scale, Tensor global_amax); -void nvfp4_fused_scale(at::Tensor block_amax, at::Tensor global_amax, at::Tensor per_block_scale, - at::Tensor target_scale, at::Tensor target_amax, int64_t tile_rows, +void nvfp4_fused_scale(Tensor block_amax, Tensor global_amax, Tensor per_block_scale, + Tensor target_scale, Tensor target_amax, int64_t tile_rows, int64_t tile_cols, int64_t rows_padded, int64_t block_len); void nvfp4_multi_tensor_fused_scale( - std::vector block_amax_list, std::vector global_amax_list, - std::vector per_block_scale_list, std::vector target_scale_list, - std::vector target_amax_list, std::vector tile_rows_list, + std::vector block_amax_list, std::vector global_amax_list, + std::vector per_block_scale_list, std::vector target_scale_list, + std::vector target_amax_list, std::vector tile_rows_list, std::vector tile_cols_list, std::vector rows_padded_list, int64_t block_len); -void nvfp4_compute_global_scale(at::Tensor global_amax, at::Tensor global_scale); +void nvfp4_compute_global_scale(Tensor global_amax, Tensor global_scale); -at::Tensor swap_first_dims(at::Tensor tensor, std::optional out = std::nullopt); +Tensor swap_first_dims(Tensor tensor, std::optional out = std::nullopt); /*************************************************************************************************** * Activations **************************************************************************************************/ /* GLU (sigmoid gate) */ -py::object glu(const at::Tensor &input, py::handle quantizer); +py::object glu(const Tensor &input, py::handle quantizer); -py::object dglu(const at::Tensor &grad, const at::Tensor &input, py::handle quantizer); +py::object dglu(const Tensor &grad, const Tensor &input, py::handle quantizer); /* GELU and variants*/ -py::object gelu(const at::Tensor &input, py::handle quantizer); +py::object gelu(const Tensor &input, py::handle quantizer); -py::object dgelu(const at::Tensor &grad, const at::Tensor &input, py::handle quantizer); +py::object dgelu(const Tensor &grad, const Tensor &input, py::handle quantizer); -py::object geglu(const at::Tensor &input, py::handle quantizer); +py::object geglu(const Tensor &input, py::handle quantizer); -py::object dgeglu(const at::Tensor &grad, const at::Tensor &input, py::handle quantizer); +py::object dgeglu(const Tensor &grad, const Tensor &input, py::handle quantizer); -py::object qgelu(const at::Tensor &input, py::handle quantizer); +py::object qgelu(const Tensor &input, py::handle quantizer); -py::object dqgelu(const at::Tensor &grad, const at::Tensor &input, py::handle quantizer); +py::object dqgelu(const Tensor &grad, const Tensor &input, py::handle quantizer); -py::object qgeglu(const at::Tensor &input, py::handle quantizer); +py::object qgeglu(const Tensor &input, py::handle quantizer); -py::object dqgeglu(const at::Tensor &grad, const at::Tensor &input, py::handle quantizer); +py::object dqgeglu(const Tensor &grad, const Tensor &input, py::handle quantizer); /* ReLU and variants*/ -py::object relu(const at::Tensor &input, py::handle quantizer); +py::object relu(const Tensor &input, py::handle quantizer); -py::object drelu(const at::Tensor &grad, const at::Tensor &input, py::handle quantizer); +py::object drelu(const Tensor &grad, const Tensor &input, py::handle quantizer); -py::object reglu(const at::Tensor &input, py::handle quantizer); +py::object reglu(const Tensor &input, py::handle quantizer); -py::object dreglu(const at::Tensor &grad, const at::Tensor &input, py::handle quantizer); +py::object dreglu(const Tensor &grad, const Tensor &input, py::handle quantizer); -py::object srelu(const at::Tensor &input, py::handle quantizer); +py::object srelu(const Tensor &input, py::handle quantizer); -py::object dsrelu(const at::Tensor &grad, const at::Tensor &input, py::handle quantizer); +py::object dsrelu(const Tensor &grad, const Tensor &input, py::handle quantizer); -py::object sreglu(const at::Tensor &input, py::handle quantizer); +py::object sreglu(const Tensor &input, py::handle quantizer); -py::object dsreglu(const at::Tensor &grad, const at::Tensor &input, py::handle quantizer); +py::object dsreglu(const Tensor &grad, const Tensor &input, py::handle quantizer); /* Silu and variants*/ -py::object silu(const at::Tensor &input, py::handle quantizer); +py::object silu(const Tensor &input, py::handle quantizer); -py::object dsilu(const at::Tensor &grad, const at::Tensor &input, py::handle quantizer); +py::object dsilu(const Tensor &grad, const Tensor &input, py::handle quantizer); -py::object swiglu(const at::Tensor &input, py::handle quantizer); +py::object swiglu(const Tensor &input, py::handle quantizer); -py::object dswiglu(const at::Tensor &grad, const at::Tensor &input, py::handle quantizer); +py::object dswiglu(const Tensor &grad, const Tensor &input, py::handle quantizer); -py::object clamped_swiglu(const at::Tensor &input, py::handle quantizer, float limit, float alpha, +py::object clamped_swiglu(const Tensor &input, py::handle quantizer, float limit, float alpha, float glu_linear_offset); -py::object clamped_dswiglu(const at::Tensor &grad, const at::Tensor &input, py::handle quantizer, +py::object clamped_dswiglu(const Tensor &grad, const Tensor &input, py::handle quantizer, float limit, float alpha, float glu_linear_offset); /*************************************************************************************************** * LayerNorm **************************************************************************************************/ -std::vector layernorm_bwd(const at::Tensor &dz, const at::Tensor &x, - const at::Tensor &mu, const at::Tensor &rsigma, - const at::Tensor &gamma, const int sm_margin, +std::vector layernorm_bwd(const Tensor &dz, const Tensor &x, + const Tensor &mu, const Tensor &rsigma, + const Tensor &gamma, const int sm_margin, const bool zero_centered_gamma); std::vector layernorm_fwd(py::handle input, py::handle weight, MaybeTensor bias, @@ -302,13 +302,13 @@ std::vector layernorm_fwd(py::handle input, py::handle weight, Maybe * RMSNorm **************************************************************************************************/ -std::vector rmsnorm_bwd(const at::Tensor &dz, const at::Tensor &x, - const at::Tensor &rsigma, const at::Tensor &gamma, +std::vector rmsnorm_bwd(const Tensor &dz, const Tensor &x, + const Tensor &rsigma, const Tensor &gamma, const int sm_margin, const bool zero_centered_gamma); -std::vector rmsnorm_bwd_add(const at::Tensor &dz, const at::Tensor &x, - const at::Tensor &add, const at::Tensor &rsigma, - const at::Tensor &gamma, const int sm_margin, +std::vector rmsnorm_bwd_add(const Tensor &dz, const Tensor &x, + const Tensor &add, const Tensor &rsigma, + const Tensor &gamma, const int sm_margin, const bool zero_centered_gamma); std::vector rmsnorm_fwd(const py::handle &input, const py::handle &weight, float eps, @@ -320,9 +320,9 @@ std::vector rmsnorm_fwd(const py::handle &input, const py::handle &w **************************************************************************************************/ // Allocates tensors all backed by a single contiguous buffer. -std::vector bulk_allocate(const std::vector> &shapes, - const std::vector &dtypes, - std::optional device = std::nullopt, +std::vector bulk_allocate(const std::vector> &shapes, + const std::vector &dtypes, + std::optional device = std::nullopt, std::optional> alignments = std::nullopt); /*************************************************************************************************** @@ -330,38 +330,38 @@ std::vector bulk_allocate(const std::vector> &sh **************************************************************************************************/ py::object create_empty_quantized_tensor(py::handle quantizer, const std::vector &shape, - at::ScalarType dtype, at::Device device, bool pin_memory); + ScalarType dtype, Device device, bool pin_memory); -py::object quantize(const at::Tensor &tensor, py::handle quantizer, const py::object &output, - std::optional noop_flag); +py::object quantize(const Tensor &tensor, py::handle quantizer, const py::object &output, + std::optional noop_flag); -py::object nvfp4_quantize_with_amax(const at::Tensor &tensor, py::handle quantizer, - const at::Tensor &rowwise_amax, - const at::Tensor &columnwise_amax); +py::object nvfp4_quantize_with_amax(const Tensor &tensor, py::handle quantizer, + const Tensor &rowwise_amax, + const Tensor &columnwise_amax); py::object dequantize(const py::handle &input, DType otype); -py::object group_quantize(const at::Tensor &tensor, py::handle quantizer, const size_t num_tensors, - std::optional first_dims, - std::optional tensor_offsets); +py::object group_quantize(const Tensor &tensor, py::handle quantizer, const size_t num_tensors, + std::optional first_dims, + std::optional tensor_offsets); -py::object nvfp4_group_quantize_with_amax(const at::Tensor &tensor, py::handle quantizer, +py::object nvfp4_group_quantize_with_amax(const Tensor &tensor, py::handle quantizer, const size_t num_tensors, - std::optional first_dims, - const at::Tensor &rowwise_amax, - const at::Tensor &columnwise_amax, - std::optional tensor_offsets); + std::optional first_dims, + const Tensor &rowwise_amax, + const Tensor &columnwise_amax, + std::optional tensor_offsets); py::object group_dequantize(const py::handle &input, DType otype); -py::object bgrad_group_quantize(const at::Tensor &tensor, py::handle quantizer, - const size_t num_tensors, std::optional first_dims, - std::optional tensor_offsets); +py::object bgrad_group_quantize(const Tensor &tensor, py::handle quantizer, + const size_t num_tensors, std::optional first_dims, + std::optional tensor_offsets); -std::vector multi_tensor_quantize(const std::vector &tensor_list, +std::vector multi_tensor_quantize(const std::vector &tensor_list, std::vector quantizer_list); -std::vector split_quantize(const at::Tensor &tensor, +std::vector split_quantize(const Tensor &tensor, const std::vector &split_sections, std::vector quantizer_list, bool disable_bulk_allocation = false); @@ -370,21 +370,21 @@ std::vector split_quantize(const at::Tensor &tensor, * Bias gradient fusions **************************************************************************************************/ -std::vector bgrad_quantize(const at::Tensor &input, py::handle py_quantizer); +std::vector bgrad_quantize(const Tensor &input, py::handle py_quantizer); -std::vector dbias_dgelu(const at::Tensor &grad_output, const at::Tensor &act_input, +std::vector dbias_dgelu(const Tensor &grad_output, const Tensor &act_input, py::handle quantizer); -std::vector dbias_dsilu(const at::Tensor &grad_output, const at::Tensor &act_input, +std::vector dbias_dsilu(const Tensor &grad_output, const Tensor &act_input, py::handle quantizer); -std::vector dbias_drelu(const at::Tensor &grad_output, const at::Tensor &act_input, +std::vector dbias_drelu(const Tensor &grad_output, const Tensor &act_input, py::handle quantizer); -std::vector dbias_dqgelu(const at::Tensor &grad_output, const at::Tensor &act_input, +std::vector dbias_dqgelu(const Tensor &grad_output, const Tensor &act_input, py::handle quantizer); -std::vector dbias_dsrelu(const at::Tensor &grad_output, const at::Tensor &act_input, +std::vector dbias_dsrelu(const Tensor &grad_output, const Tensor &act_input, py::handle quantizer); /*************************************************************************************************** @@ -392,105 +392,105 @@ std::vector dbias_dsrelu(const at::Tensor &grad_output, const at::Te **************************************************************************************************/ std::vector dropout_fwd(const py::handle &input, const float dropout_probability, - std::optional out = std::nullopt); + std::optional out = std::nullopt); -py::object dropout_bwd(const at::Tensor &grad_output, const at::Tensor &mask, +py::object dropout_bwd(const Tensor &grad_output, const Tensor &mask, const float dropout_probability, - std::optional grad_input = std::nullopt); + std::optional grad_input = std::nullopt); /*************************************************************************************************** * Softmax **************************************************************************************************/ -at::Tensor scaled_softmax_forward(at::Tensor input, float scale_factor); +Tensor scaled_softmax_forward(Tensor input, float scale_factor); -at::Tensor scaled_softmax_backward(at::Tensor output_grad_, at::Tensor softmax_results_, +Tensor scaled_softmax_backward(Tensor output_grad_, Tensor softmax_results_, float scale_factor); -at::Tensor scaled_masked_softmax_forward(at::Tensor input, at::Tensor mask, float scale_factor); +Tensor scaled_masked_softmax_forward(Tensor input, Tensor mask, float scale_factor); -at::Tensor scaled_masked_softmax_backward(at::Tensor output_grad_, at::Tensor softmax_results_, +Tensor scaled_masked_softmax_backward(Tensor output_grad_, Tensor softmax_results_, float scale_factor); -at::Tensor scaled_upper_triang_masked_softmax_forward(at::Tensor input, float scale_factor); +Tensor scaled_upper_triang_masked_softmax_forward(Tensor input, float scale_factor); -at::Tensor scaled_upper_triang_masked_softmax_backward(at::Tensor output_grads_, - at::Tensor softmax_results_, +Tensor scaled_upper_triang_masked_softmax_backward(Tensor output_grads_, + Tensor softmax_results_, float scale_factor); -at::Tensor scaled_aligned_causal_masked_softmax_forward(at::Tensor input, float scale_factor); +Tensor scaled_aligned_causal_masked_softmax_forward(Tensor input, float scale_factor); -at::Tensor scaled_aligned_causal_masked_softmax_backward(at::Tensor output_grads_, - at::Tensor softmax_results_, +Tensor scaled_aligned_causal_masked_softmax_backward(Tensor output_grads_, + Tensor softmax_results_, float scale_factor); /*************************************************************************************************** * FP8 recipe **************************************************************************************************/ -void compute_amax(const at::Tensor &tensor, at::Tensor &amax); +void compute_amax(const Tensor &tensor, Tensor &amax); -void fused_amax_and_scale_update_after_reduction(const at::Tensor &amax_reduction_buffer, - std::vector amax_histories, - std::vector scales, +void fused_amax_and_scale_update_after_reduction(const Tensor &amax_reduction_buffer, + std::vector amax_histories, + std::vector scales, const std::string &amax_compute_algo, DType fp8_dtype, float margin); // Note that the start_offset is the logical offset along the tensor dimension. // The offset in bytes is start_offset * sizeof(tensor.dtype) -void fp8_block_scaling_compute_partial_amax(const at::Tensor &tensor, at::Tensor amax, size_t h, +void fp8_block_scaling_compute_partial_amax(const Tensor &tensor, Tensor amax, size_t h, size_t w, size_t start_offset, size_t block_len); -void fp8_block_scaling_partial_cast(const at::Tensor &inp, at::Tensor out, const at::Tensor &scale, +void fp8_block_scaling_partial_cast(const Tensor &inp, Tensor out, const Tensor &scale, size_t h, size_t w, size_t start_offset, size_t block_len, const DType out_dtype); -void nvfp4_2d_compute_partial_amax(const at::Tensor &tensor, at::Tensor amax, size_t h, size_t w, +void nvfp4_2d_compute_partial_amax(const Tensor &tensor, Tensor amax, size_t h, size_t w, size_t start_offset, size_t block_len); -void nvfp4_2d_partial_cast(const at::Tensor &inp, py::handle out, const at::Tensor &scale, - const at::Tensor &global_scale, size_t h, size_t w, size_t start_offset, +void nvfp4_2d_partial_cast(const Tensor &inp, py::handle out, const Tensor &scale, + const Tensor &global_scale, size_t h, size_t w, size_t start_offset, size_t block_len); -void nvfp4_multi_tensor_2d_partial_cast(std::vector inp_list, - std::vector out_list, - std::vector scale_list, - std::vector global_scale_list, +void nvfp4_multi_tensor_2d_partial_cast(std::vector inp_list, + std::vector out_list, + std::vector scale_list, + std::vector global_scale_list, std::vector h_list, std::vector w_list, std::vector start_offset_list, int64_t block_len); -void mxfp8_scaling_compute_partial_amax(const at::Tensor &input, at::Tensor amax_rowwise, - at::Tensor amax_colwise, int rows, int cols, +void mxfp8_scaling_compute_partial_amax(const Tensor &input, Tensor amax_rowwise, + Tensor amax_colwise, int rows, int cols, size_t start_offset); -void mxfp8_scaling_partial_cast(const at::Tensor &input, at::Tensor output_rowwise, - at::Tensor output_colwise, const at::Tensor &scale_inv_rowwise, - const at::Tensor &scale_inv_colwise, int rows, int cols, +void mxfp8_scaling_partial_cast(const Tensor &input, Tensor output_rowwise, + Tensor output_colwise, const Tensor &scale_inv_rowwise, + const Tensor &scale_inv_colwise, int rows, int cols, size_t start_offset); /*************************************************************************************************** * Rotary positional embedding **************************************************************************************************/ -at::Tensor fused_rope_forward(const at::Tensor &input, const at::Tensor &freqs, - const std::optional start_positions, +Tensor fused_rope_forward(const Tensor &input, const Tensor &freqs, + const std::optional start_positions, const NVTE_QKV_Format qkv_format, const bool interleaved, - const std::optional cu_seqlens, const int cp_size, + const std::optional cu_seqlens, const int cp_size, const int cp_rank); -at::Tensor fused_rope_backward(const at::Tensor &output_grads, const at::Tensor &freqs, - const std::optional start_positions, +Tensor fused_rope_backward(const Tensor &output_grads, const Tensor &freqs, + const std::optional start_positions, const NVTE_QKV_Format qkv_format, const bool interleaved, - const std::optional cu_seqlens, const int cp_size, + const std::optional cu_seqlens, const int cp_size, const int cp_rank); -std::tuple fused_qkv_rope_forward( - const at::Tensor &qkv_input, const at::Tensor &q_freqs, const at::Tensor &k_freqs, - const std::optional start_positions, const std::vector &qkv_split_arg_list, +std::tuple fused_qkv_rope_forward( + const Tensor &qkv_input, const Tensor &q_freqs, const Tensor &k_freqs, + const std::optional start_positions, const std::vector &qkv_split_arg_list, const NVTE_QKV_Format qkv_format, const bool interleaved, const int cp_size, const int cp_rank); -at::Tensor fused_qkv_rope_backward(const at::Tensor &q_grad_out, const at::Tensor &k_grad_out, - const at::Tensor &v_grad_out, const at::Tensor &q_freqs, - const at::Tensor &k_freqs, +Tensor fused_qkv_rope_backward(const Tensor &q_grad_out, const Tensor &k_grad_out, + const Tensor &v_grad_out, const Tensor &q_freqs, + const Tensor &k_freqs, const std::vector &qkv_split_arg_list, const NVTE_QKV_Format qkv_format, const bool interleaved, const int cp_size, const int cp_rank); @@ -503,14 +503,14 @@ size_t get_cublasLt_version(); size_t get_cudnn_version(); -at::Tensor splits_to_offsets(const at::Tensor &first_dims, int64_t logical_last_dim); -std::tuple> splits_to_offsets_multi( - const at::Tensor &split_sizes, const c10::Device &device, const std::vector &strides, - const std::vector &include_leading_zero, const std::vector &dtypes, +Tensor splits_to_offsets(const Tensor &first_dims, int64_t logical_last_dim); +std::tuple> splits_to_offsets_multi( + const Tensor &split_sizes, const Device &device, const std::vector &strides, + const std::vector &include_leading_zero, const std::vector &dtypes, bool bulk_allocate_outputs); -at::Tensor copy_data_ptrs_to_device(const std::vector &tensors, - const c10::Device &device); +Tensor copy_data_ptrs_to_device(const std::vector &tensors, + const Device &device); /*************************************************************************************************** * Experimental helpers for the fused grouped MLP @@ -527,9 +527,9 @@ namespace grouped_mlp_experimental { // device. All tensors must share a uniform shape and `swizzle_type` // must be one of "mxfp8_rowwise", "mxfp8_columnwise", or "nvfp4". // Returns {data_ptrs_device, scale_ptrs_device, swizzled_scales_buffer}. -std::tuple swizzle_scales_and_pack_ptrs_for_discrete_weights( - const std::vector &data_tensors, const std::vector &scale_tensors, - const std::string &swizzle_type, const c10::Device &device); +std::tuple swizzle_scales_and_pack_ptrs_for_discrete_weights( + const std::vector &data_tensors, const std::vector &scale_tensors, + const std::string &swizzle_type, const Device &device); } // namespace grouped_mlp_experimental @@ -537,98 +537,98 @@ std::tuple swizzle_scales_and_pack_ptrs_for_ * Support THD format for Context Parallel **************************************************************************************************/ -at::Tensor thd_read_half_tensor(const at::Tensor &tensor, const at::Tensor &cu_seqlens, +Tensor thd_read_half_tensor(const Tensor &tensor, const Tensor &cu_seqlens, int half_idx); -void thd_second_half_lse_correction(at::Tensor lse, const at::Tensor &lse_per_step, - const at::Tensor &cu_seqlens, bool lse_packed); +void thd_second_half_lse_correction(Tensor lse, const Tensor &lse_per_step, + const Tensor &cu_seqlens, bool lse_packed); -at::Tensor thd_read_second_half_lse(const at::Tensor &lse, const at::Tensor &cu_seqlens, +Tensor thd_read_second_half_lse(const Tensor &lse, const Tensor &cu_seqlens, bool lse_packed, int second_half_lse_seqlen); -void thd_out_correction(at::Tensor out, const at::Tensor &out_per_step, const at::Tensor &lse, - const at::Tensor &lse_per_step, const at::Tensor &cu_seqlens, +void thd_out_correction(Tensor out, const Tensor &out_per_step, const Tensor &lse, + const Tensor &lse_per_step, const Tensor &cu_seqlens, bool only_second_half, bool lse_packed); -void thd_grad_correction(at::Tensor grad, const at::Tensor &grad_per_step, - const at::Tensor &cu_seqlens, const std::string &first_half, +void thd_grad_correction(Tensor grad, const Tensor &grad_per_step, + const Tensor &cu_seqlens, const std::string &first_half, const std::string &second_half); -at::Tensor thd_get_partitioned_indices(const at::Tensor &cu_seqlens, int total_tokens, +Tensor thd_get_partitioned_indices(const Tensor &cu_seqlens, int total_tokens, int world_size, int rank); /*************************************************************************************************** * multi_tensor_* kernels **************************************************************************************************/ -void multi_tensor_scale_cuda(int chunk_size, at::Tensor noop_flag, - std::vector> tensor_lists, float scale); +void multi_tensor_scale_cuda(int chunk_size, Tensor noop_flag, + std::vector> tensor_lists, float scale); -void multi_tensor_scale_tensor_cuda(int chunk_size, at::Tensor is_infinite, - std::vector> tensor_lists, - at::Tensor scale); +void multi_tensor_scale_tensor_cuda(int chunk_size, Tensor is_infinite, + std::vector> tensor_lists, + Tensor scale); -std::tuple multi_tensor_l2norm_cuda( - int chunk_size, at::Tensor noop_flag, std::vector> tensor_lists, - at::optional per_tensor_python); +std::tuple multi_tensor_l2norm_cuda( + int chunk_size, Tensor noop_flag, std::vector> tensor_lists, + std::optional per_tensor_python); -std::tuple multi_tensor_unscale_l2norm_cuda( - int chunk_size, at::Tensor noop_flag, std::vector> tensor_lists, - at::Tensor inv_scale, at::optional per_tensor_python); +std::tuple multi_tensor_unscale_l2norm_cuda( + int chunk_size, Tensor noop_flag, std::vector> tensor_lists, + Tensor inv_scale, std::optional per_tensor_python); -void multi_tensor_adam_cuda(int chunk_size, at::Tensor noop_flag, - std::vector> tensor_lists, const float lr, +void multi_tensor_adam_cuda(int chunk_size, Tensor noop_flag, + std::vector> tensor_lists, const float lr, const float beta1, const float beta2, const float epsilon, const int step, const int mode, const int bias_correction, const float weight_decay); -void multi_tensor_adam_param_remainder_cuda(int chunk_size, at::Tensor noop_flag, - std::vector> tensor_lists, +void multi_tensor_adam_param_remainder_cuda(int chunk_size, Tensor noop_flag, + std::vector> tensor_lists, const float lr, const float beta1, const float beta2, const float epsilon, const int step, const int mode, const int bias_correction, const float weight_decay); -void multi_tensor_adam_fp8_cuda(int chunk_size, at::Tensor noop_flag, - std::vector> tensor_lists, const float lr, +void multi_tensor_adam_fp8_cuda(int chunk_size, Tensor noop_flag, + std::vector> tensor_lists, const float lr, const float beta1, const float beta2, const float epsilon, const int step, const int mode, const int bias_correction, const float weight_decay, DType fp8_dtype); -void multi_tensor_adam_capturable_cuda(int chunk_size, at::Tensor noop_flag, - std::vector> tensor_lists, - at::Tensor lr, const float beta1, const float beta2, - const float epsilon, at::Tensor step, const int mode, +void multi_tensor_adam_capturable_cuda(int chunk_size, Tensor noop_flag, + std::vector> tensor_lists, + Tensor lr, const float beta1, const float beta2, + const float epsilon, Tensor step, const int mode, const int bias_correction, const float weight_decay, - at::Tensor inv_scale); + Tensor inv_scale); -void multi_tensor_adam_capturable_master_cuda(int chunk_size, at::Tensor noop_flag, - std::vector> tensor_lists, - at::Tensor lr, const float beta1, const float beta2, - const float epsilon, at::Tensor step, const int mode, +void multi_tensor_adam_capturable_master_cuda(int chunk_size, Tensor noop_flag, + std::vector> tensor_lists, + Tensor lr, const float beta1, const float beta2, + const float epsilon, Tensor step, const int mode, const int bias_correction, const float weight_decay, - at::Tensor inv_scale); + Tensor inv_scale); -void multi_tensor_sgd_cuda(int chunk_size, at::Tensor noop_flag, - std::vector> tensor_lists, float wd, +void multi_tensor_sgd_cuda(int chunk_size, Tensor noop_flag, + std::vector> tensor_lists, float wd, float momentum, float dampening, float lr, bool nesterov, bool first_run, bool wd_after_momentum, float scale); void multi_tensor_compute_scale_and_scale_inv_cuda( - int chunk_size, at::Tensor noop_flag, std::vector> tensor_lists, + int chunk_size, Tensor noop_flag, std::vector> tensor_lists, float max_fp8, bool force_pow_2_scales, float epsilon); void multi_tensor_compute_scale_inv_e8m0_cuda(int chunk_size, const py::object &dummy, - std::vector> tensor_lists); + std::vector> tensor_lists); /*************************************************************************************************** * padding **************************************************************************************************/ -void fused_multi_row_padding(at::Tensor input, at::Tensor output, +void fused_multi_row_padding(Tensor input, Tensor output, std::vector input_row_list, std::vector padded_input_row_list); -void fused_multi_row_unpadding(at::Tensor input, at::Tensor output, +void fused_multi_row_unpadding(Tensor input, Tensor output, std::vector input_row_list, std::vector unpadded_input_row_list); @@ -651,13 +651,13 @@ void grouped_swizzle_for_gemm(py::handle &tensor, bool rowwise, bool columnwise) * NVSHMEM APIs **************************************************************************************************/ -void init_nvshmem_backend(c10d::ProcessGroup *process_group); +void init_nvshmem_backend(ProcessGroup *process_group); -at::Tensor create_nvshmem_tensor(const std::vector &shape, c10::ScalarType dtype); +Tensor create_nvshmem_tensor(const std::vector &shape, ScalarType dtype); -void nvshmem_send_on_current_stream(at::Tensor src, at::Tensor dst, int peer, at::Tensor signal); +void nvshmem_send_on_current_stream(Tensor src, Tensor dst, int peer, Tensor signal); -void nvshmem_wait_on_current_stream(at::Tensor signal, const std::string &wait_kind); +void nvshmem_wait_on_current_stream(Tensor signal, const std::string &wait_kind); void nvshmem_finalize(); @@ -665,8 +665,8 @@ void nvshmem_finalize(); * Comm+GEMM Overlap Wrappers **************************************************************************************************/ -void bulk_overlap_ag_with_external_gemm(CommOverlap &allgather_communicator, at::Stream send_stream, - at::Stream recv_stream); +void bulk_overlap_ag_with_external_gemm(CommOverlap &allgather_communicator, Stream send_stream, + Stream recv_stream); /*************************************************************************************************** * Newton-Schulz (cuSolverMp) @@ -676,7 +676,7 @@ int64_t cusolvermp_ctx_create(int64_t nccl_comm_ptr, int nranks, int rank); void cusolvermp_ctx_destroy(int64_t ctx_ptr); -void newton_schulz(int64_t ctx_ptr, int64_t m, int64_t n, at::Tensor x, int64_t num_iterations, +void newton_schulz(int64_t ctx_ptr, int64_t m, int64_t n, Tensor x, int64_t num_iterations, std::vector coefficients); } // namespace transformer_engine::pytorch @@ -685,7 +685,7 @@ void newton_schulz(int64_t ctx_ptr, int64_t m, int64_t n, at::Tensor x, int64_t * Comm+GEMM Overlap Wrappers **************************************************************************************************/ -class CommOverlapHelper : torch::CustomClassHolder { +class CommOverlapHelper : transformer_engine::pytorch::CustomClassHolder { public: // Shared ownership of an ncclComm_t. The deleter calls ncclCommDestroy when // the last reference (held by the helper and/or any CommOverlap consumers) @@ -696,7 +696,7 @@ class CommOverlapHelper : torch::CustomClassHolder { private: bool initialized{false}; bool backend_is_nccl{false}; - std::map torch_pgs; + std::map torch_pgs; std::map nccl_comms; public: @@ -709,8 +709,8 @@ class CommOverlapHelper : torch::CustomClassHolder { CommOverlapHelper(); - CommOverlapHelper(c10d::ProcessGroup *world_group, - std::optional intra_node_group); + CommOverlapHelper(transformer_engine::pytorch::ProcessGroup *world_group, + std::optional intra_node_group); ~CommOverlapHelper(); @@ -722,14 +722,14 @@ class CommOverlapHelper : torch::CustomClassHolder { NcclCommSharedPtr get_nccl_comm(std::string comm_name); }; -class CommOverlap : torch::CustomClassHolder, public transformer_engine::CommOverlapBase { +class CommOverlap : transformer_engine::pytorch::CustomClassHolder, public transformer_engine::CommOverlapBase { private: // Keeps the cuBLASMp NCCL communicator alive for the lifetime of this // instance, independent of the CommOverlapHelper that created it. CommOverlapHelper::NcclCommSharedPtr _nccl_comm; public: - CommOverlap(const std::vector &buffer_shape, at::ScalarType buffer_dtype, + CommOverlap(const std::vector &buffer_shape, transformer_engine::pytorch::ScalarType buffer_dtype, CommOverlapHelper *helper, int tp_size, int num_splits = 4, int num_max_streams = NVTE_COMM_OVERLAP_MAX_STREAMS, int comm_cga_size = 2, int gemm_priority = 0, int comm_priority = 0, int num_comm_sm = 16, @@ -742,29 +742,29 @@ class CommOverlap : torch::CustomClassHolder, public transformer_engine::CommOve // (including those captured in CUDA graphs) avoid the unsafe lazy paths. CommOverlap(CommOverlapHelper *helper, int tp_rank, int tp_size, transformer_engine::CommOverlapType comm_type, - const std::vector &buffer_shape, at::ScalarType buffer_dtype, + const std::vector &buffer_shape, transformer_engine::pytorch::ScalarType buffer_dtype, int num_comm_sm = 16, bool atomic_gemm = false); ~CommOverlap() {} using transformer_engine::CommOverlapCore::copy_into_buffer; - void copy_into_buffer(const at::Tensor &input, bool local_chunk = false); + void copy_into_buffer(const transformer_engine::pytorch::Tensor &input, bool local_chunk = false); - at::Tensor get_buffer(bool local_chunk = false, + transformer_engine::pytorch::Tensor get_buffer(bool local_chunk = false, std::optional> shape = std::nullopt); - std::pair get_communication_stream(); + std::pair get_communication_stream(); }; // CommOverlap -class CommOverlapP2P : torch::CustomClassHolder, public transformer_engine::CommOverlapP2PBase { +class CommOverlapP2P : transformer_engine::pytorch::CustomClassHolder, public transformer_engine::CommOverlapP2PBase { private: // Keeps the cuBLASMp NCCL communicator alive for the lifetime of this // instance, independent of the CommOverlapHelper that created it. CommOverlapHelper::NcclCommSharedPtr _nccl_comm; public: - CommOverlapP2P(const std::vector &buffer_shape, at::ScalarType buffer_dtype, + CommOverlapP2P(const std::vector &buffer_shape, transformer_engine::pytorch::ScalarType buffer_dtype, CommOverlapHelper *helper, int tp_size, transformer_engine::CommOverlapType comm_type, int num_max_streams = NVTE_COMM_OVERLAP_MAX_STREAMS, int comm_cga_size = 1, @@ -775,18 +775,18 @@ class CommOverlapP2P : torch::CustomClassHolder, public transformer_engine::Comm // cuBLASMp variant. See CommOverlap for the `comm_type`/buffer args. CommOverlapP2P(CommOverlapHelper *helper, int tp_rank, int tp_size, transformer_engine::CommOverlapType comm_type, - const std::vector &buffer_shape, at::ScalarType buffer_dtype, + const std::vector &buffer_shape, transformer_engine::pytorch::ScalarType buffer_dtype, int num_comm_sm = 1, bool atomic_gemm = false); ~CommOverlapP2P() {} using transformer_engine::CommOverlapP2PBase::copy_into_buffer; - void copy_into_buffer(const at::Tensor &input, bool local_chunk = false); + void copy_into_buffer(const transformer_engine::pytorch::Tensor &input, bool local_chunk = false); - at::Tensor get_buffer(bool local_chunk = false, + transformer_engine::pytorch::Tensor get_buffer(bool local_chunk = false, std::optional> shape = std::nullopt); - std::pair get_communication_stream(); + std::pair get_communication_stream(); }; // CommOverlapP2P diff --git a/transformer_engine/pytorch/csrc/extensions/activation.cpp b/transformer_engine/pytorch/csrc/extensions/activation.cpp index 58a8f84f85..e850f20efe 100644 --- a/transformer_engine/pytorch/csrc/extensions/activation.cpp +++ b/transformer_engine/pytorch/csrc/extensions/activation.cpp @@ -15,7 +15,7 @@ using FuncType = void(const NVTETensor, NVTETensor, cudaStream_t); using DFuncType = void(const NVTETensor, const NVTETensor, NVTETensor, cudaStream_t); template -py::object activation_helper(const at::Tensor& input, py::handle quantizer, int shape_divisor = 1, +py::object activation_helper(const Tensor& input, py::handle quantizer, int shape_divisor = 1, Args&&... args) { init_extension(); @@ -52,7 +52,7 @@ py::object activation_helper(const at::Tensor& input, py::handle quantizer, int } // Perform compute - auto stream = at::cuda::getCurrentCUDAStream(); + auto stream = getCurrentCUDAStream(); switch (impl) { case Impl::UNFUSED: // Compute activation in high precision, then quantize @@ -126,7 +126,7 @@ py::object activation_helper(const at::Tensor& input, py::handle quantizer, int } template -py::object dactivation_helper(const at::Tensor& grad_output, const at::Tensor& input, +py::object dactivation_helper(const Tensor& grad_output, const Tensor& input, py::handle quantizer, Args&&... args) { init_extension(); @@ -165,7 +165,7 @@ py::object dactivation_helper(const at::Tensor& grad_output, const at::Tensor& i } // Perform compute - auto stream = at::cuda::getCurrentCUDAStream(); + auto stream = getCurrentCUDAStream(); switch (impl) { case Impl::UNFUSED: // Compute activation backward in high precision, then quantize @@ -240,103 +240,103 @@ py::object dactivation_helper(const at::Tensor& grad_output, const at::Tensor& i } // namespace /* GELU and variants */ -py::object gelu(const at::Tensor& input, py::handle quantizer) { +py::object gelu(const Tensor& input, py::handle quantizer) { return activation_helper(input, quantizer); } -py::object dgelu(const at::Tensor& grad, const at::Tensor& input, py::handle quantizer) { +py::object dgelu(const Tensor& grad, const Tensor& input, py::handle quantizer) { return dactivation_helper(grad, input, quantizer); } -py::object glu(const at::Tensor& input, py::handle quantizer) { +py::object glu(const Tensor& input, py::handle quantizer) { return activation_helper(input, quantizer, 2); } -py::object dglu(const at::Tensor& grad, const at::Tensor& input, py::handle quantizer) { +py::object dglu(const Tensor& grad, const Tensor& input, py::handle quantizer) { return dactivation_helper(grad, input, quantizer); } -py::object geglu(const at::Tensor& input, py::handle quantizer) { +py::object geglu(const Tensor& input, py::handle quantizer) { return activation_helper(input, quantizer, 2); } -py::object dgeglu(const at::Tensor& grad, const at::Tensor& input, py::handle quantizer) { +py::object dgeglu(const Tensor& grad, const Tensor& input, py::handle quantizer) { return dactivation_helper(grad, input, quantizer); } -py::object qgelu(const at::Tensor& input, py::handle quantizer) { +py::object qgelu(const Tensor& input, py::handle quantizer) { return activation_helper(input, quantizer); } -py::object dqgelu(const at::Tensor& grad, const at::Tensor& input, py::handle quantizer) { +py::object dqgelu(const Tensor& grad, const Tensor& input, py::handle quantizer) { return dactivation_helper(grad, input, quantizer); } -py::object qgeglu(const at::Tensor& input, py::handle quantizer) { +py::object qgeglu(const Tensor& input, py::handle quantizer) { return activation_helper(input, quantizer, 2); } -py::object dqgeglu(const at::Tensor& grad, const at::Tensor& input, py::handle quantizer) { +py::object dqgeglu(const Tensor& grad, const Tensor& input, py::handle quantizer) { return dactivation_helper(grad, input, quantizer); } /* ReLU and variants */ -py::object relu(const at::Tensor& input, py::handle quantizer) { +py::object relu(const Tensor& input, py::handle quantizer) { return activation_helper(input, quantizer); } -py::object drelu(const at::Tensor& grad, const at::Tensor& input, py::handle quantizer) { +py::object drelu(const Tensor& grad, const Tensor& input, py::handle quantizer) { return dactivation_helper(grad, input, quantizer); } -py::object reglu(const at::Tensor& input, py::handle quantizer) { +py::object reglu(const Tensor& input, py::handle quantizer) { return activation_helper(input, quantizer, 2); } -py::object dreglu(const at::Tensor& grad, const at::Tensor& input, py::handle quantizer) { +py::object dreglu(const Tensor& grad, const Tensor& input, py::handle quantizer) { return dactivation_helper(grad, input, quantizer); } -py::object srelu(const at::Tensor& input, py::handle quantizer) { +py::object srelu(const Tensor& input, py::handle quantizer) { return activation_helper(input, quantizer); } -py::object dsrelu(const at::Tensor& grad, const at::Tensor& input, py::handle quantizer) { +py::object dsrelu(const Tensor& grad, const Tensor& input, py::handle quantizer) { return dactivation_helper(grad, input, quantizer); } -py::object sreglu(const at::Tensor& input, py::handle quantizer) { +py::object sreglu(const Tensor& input, py::handle quantizer) { return activation_helper(input, quantizer, 2); } -py::object dsreglu(const at::Tensor& grad, const at::Tensor& input, py::handle quantizer) { +py::object dsreglu(const Tensor& grad, const Tensor& input, py::handle quantizer) { return dactivation_helper(grad, input, quantizer); } /* Silu and variants */ -py::object silu(const at::Tensor& input, py::handle quantizer) { +py::object silu(const Tensor& input, py::handle quantizer) { return activation_helper(input, quantizer); } -py::object dsilu(const at::Tensor& grad, const at::Tensor& input, py::handle quantizer) { +py::object dsilu(const Tensor& grad, const Tensor& input, py::handle quantizer) { return dactivation_helper(grad, input, quantizer); } -py::object swiglu(const at::Tensor& input, py::handle quantizer) { +py::object swiglu(const Tensor& input, py::handle quantizer) { return activation_helper(input, quantizer, 2); } -py::object dswiglu(const at::Tensor& grad, const at::Tensor& input, py::handle quantizer) { +py::object dswiglu(const Tensor& grad, const Tensor& input, py::handle quantizer) { return dactivation_helper(grad, input, quantizer); } /* clamped functions */ -py::object clamped_swiglu(const at::Tensor& input, py::handle quantizer, float limit, float alpha, +py::object clamped_swiglu(const Tensor& input, py::handle quantizer, float limit, float alpha, float glu_linear_offset) { return activation_helper(input, quantizer, 2, limit, alpha, glu_linear_offset); } -py::object clamped_dswiglu(const at::Tensor& grad, const at::Tensor& input, py::handle quantizer, +py::object clamped_dswiglu(const Tensor& grad, const Tensor& input, py::handle quantizer, float limit, float alpha, float glu_linear_offset) { return dactivation_helper(grad, input, quantizer, limit, alpha, glu_linear_offset); diff --git a/transformer_engine/pytorch/csrc/extensions/allocate.cpp b/transformer_engine/pytorch/csrc/extensions/allocate.cpp index 62ed7f3739..ca126bc917 100644 --- a/transformer_engine/pytorch/csrc/extensions/allocate.cpp +++ b/transformer_engine/pytorch/csrc/extensions/allocate.cpp @@ -22,9 +22,9 @@ namespace pytorch { * Stream usage is not recorded, so there may be race conditions if * compute is performed on multiple streams. */ -std::vector bulk_allocate(const std::vector> &shapes, - const std::vector &dtypes, - std::optional device, +std::vector bulk_allocate(const std::vector> &shapes, + const std::vector &dtypes, + std::optional device, std::optional> alignments) { // Check shapes and dtypes const size_t n = shapes.size(); @@ -37,13 +37,13 @@ std::vector bulk_allocate(const std::vector> &sh // Set defaults for optional arguments if (!device) { - device = c10::Device(c10::kCUDA); + device = Device(kCUDA); } if (!alignments) { alignments = std::vector{}; alignments->reserve(n); for (const auto &dtype : dtypes) { - alignments->push_back(c10::elementSize(dtype)); + alignments->push_back(elementSize(dtype)); } } @@ -53,7 +53,7 @@ std::vector bulk_allocate(const std::vector> &sh size_t base_byte_size = 0; size_t base_alignment = 1; for (size_t i = 0; i < n; ++i) { - byte_sizes[i] = product(shapes[i]) * at::elementSize(dtypes[i]); + byte_sizes[i] = product(shapes[i]) * elementSize(dtypes[i]); offsets[i] = roundup(base_byte_size, (*alignments)[i]); base_byte_size = offsets[i] + byte_sizes[i]; base_alignment = std::max(base_alignment, (*alignments)[i]); @@ -64,14 +64,14 @@ std::vector bulk_allocate(const std::vector> &sh } // Allocate base buffer - auto base_buffer = std::make_shared( - at::empty({static_cast(base_byte_size)}, at::device(*device).dtype(torch::kUInt8))); + auto base_buffer = std::make_shared( + empty({static_cast(base_byte_size)}, TensorOptions().device(*device).dtype(kUInt8))); uint8_t *base_ptr = base_buffer->data_ptr(); base_ptr = reinterpret_cast(roundup(reinterpret_cast(base_ptr), base_alignment)); // Create views into base buffer - std::vector out; + std::vector out; out.reserve(n); std::vector shape_int64; for (size_t i = 0; i < n; ++i) { @@ -81,12 +81,12 @@ std::vector bulk_allocate(const std::vector> &sh // empty tensor. Passing a null pointer fails because it checks // that the pointer is on GPU. Passing a non-null pointer can // cause bugs in TE kernels. - out.emplace_back(at::empty(shape_int64, at::device(*device).dtype(dtypes[i]))); + out.emplace_back(empty(shape_int64, TensorOptions().device(*device).dtype(dtypes[i]))); } else { // Construct tensor with custom deleter to keep base buffer alive - out.emplace_back(at::from_blob( + out.emplace_back(from_blob( base_ptr + offsets[i], shape_int64, [base_buffer](void *) {}, - at::device(*device).dtype(dtypes[i]))); + TensorOptions().device(*device).dtype(dtypes[i]))); } } return out; diff --git a/transformer_engine/pytorch/csrc/extensions/apply_rope.cpp b/transformer_engine/pytorch/csrc/extensions/apply_rope.cpp index 4392fa4b43..12a8a1fa6b 100644 --- a/transformer_engine/pytorch/csrc/extensions/apply_rope.cpp +++ b/transformer_engine/pytorch/csrc/extensions/apply_rope.cpp @@ -9,20 +9,20 @@ namespace transformer_engine::pytorch { -at::Tensor fused_rope_forward(const at::Tensor &input, const at::Tensor &freqs, - const std::optional start_positions, +Tensor fused_rope_forward(const Tensor &input, const Tensor &freqs, + const std::optional start_positions, const NVTE_QKV_Format qkv_format, const bool interleaved, - const std::optional cu_seqlens, const int cp_size, + const std::optional cu_seqlens, const int cp_size, const int cp_rank) { TORCH_CHECK(freqs.dim() == 4, "expected 4D tensor"); TORCH_CHECK(freqs.size(1) == 1 && freqs.size(2) == 1, "expected the second and third dims of the freqs tensor equal 1"); - TORCH_CHECK(freqs.scalar_type() == at::ScalarType::Float, + TORCH_CHECK(freqs.scalar_type() == ScalarType::Float, "Dtype of the freqs tensor must be float"); // output - auto act_options = at::TensorOptions().dtype(input.scalar_type()).device(input.device()); - auto output = at::empty(input.sizes(), act_options); + auto act_options = TensorOptions().dtype(input.scalar_type()).device(input.device()); + auto output = empty(input.sizes(), act_options); auto input_cu = makeTransformerEngineTensor(input); auto freqs_cu = makeTransformerEngineTensor(freqs); @@ -64,7 +64,7 @@ at::Tensor fused_rope_forward(const at::Tensor &input, const at::Tensor &freqs, nvte_fused_rope_forward(input_cu.data(), cu_seqlens_cu.data(), freqs_cu.data(), start_positions_cu.data(), output_cu.data(), qkv_format, interleaved, cp_size, cp_rank, max_s, b, h, d, d2, stride_t, /*stride_b=*/0, - stride_h, stride_d, at::cuda::getCurrentCUDAStream()); + stride_h, stride_d, getCurrentCUDAStream()); return output; } @@ -98,38 +98,38 @@ at::Tensor fused_rope_forward(const at::Tensor &input, const at::Tensor &freqs, nvte_fused_rope_forward(input_cu.data(), cu_seqlens_cu.data(), freqs_cu.data(), start_positions_cu.data(), output_cu.data(), qkv_format, interleaved, cp_size, cp_rank, s, b, h, d, d2, stride_s, stride_b, stride_h, stride_d, - at::cuda::getCurrentCUDAStream()); + getCurrentCUDAStream()); return output; } -std::tuple fused_qkv_rope_forward( - const at::Tensor &qkv_input, const at::Tensor &q_freqs, const at::Tensor &k_freqs, - const std::optional start_positions, const std::vector &qkv_split_arg_list, +std::tuple fused_qkv_rope_forward( + const Tensor &qkv_input, const Tensor &q_freqs, const Tensor &k_freqs, + const std::optional start_positions, const std::vector &qkv_split_arg_list, const NVTE_QKV_Format qkv_format, const bool interleaved, const int cp_size, const int cp_rank) { TORCH_CHECK(q_freqs.dim() == 4, "expected 4D tensor"); TORCH_CHECK(q_freqs.size(1) == 1 && q_freqs.size(2) == 1, "expected the second and third dims of the freqs tensor equal 1"); - TORCH_CHECK(q_freqs.scalar_type() == at::ScalarType::Float, + TORCH_CHECK(q_freqs.scalar_type() == ScalarType::Float, "Dtype of the freqs tensor must be float"); TORCH_CHECK(k_freqs.dim() == 4, "expected 4D tensor"); TORCH_CHECK(k_freqs.size(1) == 1 && k_freqs.size(2) == 1, "expected the second and third dims of the freqs tensor equal 1"); - TORCH_CHECK(k_freqs.scalar_type() == at::ScalarType::Float, + TORCH_CHECK(k_freqs.scalar_type() == ScalarType::Float, "Dtype of the freqs tensor must be float"); // output - auto act_options = at::TensorOptions().dtype(qkv_input.scalar_type()).device(qkv_input.device()); + auto act_options = TensorOptions().dtype(qkv_input.scalar_type()).device(qkv_input.device()); auto q_out_size = qkv_input.sizes().vec(); q_out_size[2] = q_out_size[2] * qkv_split_arg_list[0] / qkv_split_arg_list[1]; q_out_size[3] = qkv_split_arg_list[1]; - auto q_out = at::empty(q_out_size, act_options); + auto q_out = empty(q_out_size, act_options); auto k_out_size = qkv_input.sizes().vec(); k_out_size[3] = qkv_split_arg_list[1]; - auto k_out = at::empty(k_out_size, act_options); + auto k_out = empty(k_out_size, act_options); auto v_out_size = qkv_input.sizes().vec(); v_out_size[3] = qkv_split_arg_list[2]; - auto v_out = at::empty(v_out_size, act_options); + auto v_out = empty(v_out_size, act_options); auto qkv_cu = makeTransformerEngineTensor(qkv_input); auto q_freqs_cu = makeTransformerEngineTensor(q_freqs); @@ -157,25 +157,25 @@ std::tuple fused_qkv_rope_forward( start_positions_cu.data(), q_out_cu.data(), k_out_cu.data(), v_out_cu.data(), qkv_format, interleaved, cp_size, cp_rank, s, b, h, d, d2, qkv_split_arg_list[0], qkv_split_arg_list[1], - qkv_split_arg_list[2], at::cuda::getCurrentCUDAStream()); + qkv_split_arg_list[2], getCurrentCUDAStream()); return std::make_tuple(q_out, k_out, v_out); } -at::Tensor fused_rope_backward(const at::Tensor &output_grads, const at::Tensor &freqs, - const std::optional start_positions, +Tensor fused_rope_backward(const Tensor &output_grads, const Tensor &freqs, + const std::optional start_positions, const NVTE_QKV_Format qkv_format, const bool interleaved, - const std::optional cu_seqlens, const int cp_size, + const std::optional cu_seqlens, const int cp_size, const int cp_rank) { TORCH_CHECK(freqs.dim() == 4, "expected 4D tensor"); TORCH_CHECK(freqs.size(1) == 1 && freqs.size(2) == 1, "expected the second and third dims of the freqs tensor equal 1"); - TORCH_CHECK(freqs.scalar_type() == at::ScalarType::Float, + TORCH_CHECK(freqs.scalar_type() == ScalarType::Float, "Dtype of the freqs tensor must be float"); auto act_options = - at::TensorOptions().dtype(output_grads.scalar_type()).device(output_grads.device()); - auto input_grads = at::empty(output_grads.sizes(), act_options); + TensorOptions().dtype(output_grads.scalar_type()).device(output_grads.device()); + auto input_grads = empty(output_grads.sizes(), act_options); auto output_grads_cu = makeTransformerEngineTensor(output_grads); auto freqs_cu = makeTransformerEngineTensor(freqs); @@ -217,7 +217,7 @@ at::Tensor fused_rope_backward(const at::Tensor &output_grads, const at::Tensor nvte_fused_rope_backward(output_grads_cu.data(), cu_seqlens_cu.data(), freqs_cu.data(), start_positions_cu.data(), input_grads_cu.data(), qkv_format, interleaved, cp_size, cp_rank, max_s, b, h, d, d2, stride_t, - /*stride_b=*/0, stride_h, stride_d, at::cuda::getCurrentCUDAStream()); + /*stride_b=*/0, stride_h, stride_d, getCurrentCUDAStream()); return input_grads; } @@ -255,26 +255,26 @@ at::Tensor fused_rope_backward(const at::Tensor &output_grads, const at::Tensor nvte_fused_rope_backward(output_grads_cu.data(), cu_seqlens_cu.data(), freqs_cu.data(), start_positions_cu.data(), input_grads_cu.data(), qkv_format, interleaved, cp_size, cp_rank, s, b, h, d, d2, stride_s, stride_b, - stride_h, stride_d, at::cuda::getCurrentCUDAStream()); + stride_h, stride_d, getCurrentCUDAStream()); return input_grads; } -at::Tensor fused_qkv_rope_backward(const at::Tensor &q_grad_out, const at::Tensor &k_grad_out, - const at::Tensor &v_grad_out, const at::Tensor &q_freqs, - const at::Tensor &k_freqs, +Tensor fused_qkv_rope_backward(const Tensor &q_grad_out, const Tensor &k_grad_out, + const Tensor &v_grad_out, const Tensor &q_freqs, + const Tensor &k_freqs, const std::vector &qkv_split_arg_list, const NVTE_QKV_Format qkv_format, const bool interleaved, const int cp_size, const int cp_rank) { auto act_options = - at::TensorOptions().dtype(q_grad_out.scalar_type()).device(q_grad_out.device()); + TensorOptions().dtype(q_grad_out.scalar_type()).device(q_grad_out.device()); auto qkv_grad_size = q_grad_out.sizes().vec(); auto total_hd = (q_grad_out.size(2) + k_grad_out.size(2) + v_grad_out.size(2)) * q_grad_out.size(3); auto total_d = qkv_split_arg_list[0] + qkv_split_arg_list[1] + qkv_split_arg_list[2]; qkv_grad_size[2] = total_hd / total_d; qkv_grad_size[3] = total_d; - auto qkv_grad_input = at::empty(qkv_grad_size, act_options); + auto qkv_grad_input = empty(qkv_grad_size, act_options); const bool is_sbhd = qkv_format == NVTE_QKV_Format::NVTE_SBHD; const int s = is_sbhd ? q_grad_out.size(0) : q_grad_out.size(1); const int b = is_sbhd ? q_grad_out.size(1) : q_grad_out.size(0); @@ -293,7 +293,7 @@ at::Tensor fused_qkv_rope_backward(const at::Tensor &q_grad_out, const at::Tenso q_freqs_cu.data(), k_freqs_cu.data(), qkv_grad_cu.data(), qkv_format, interleaved, cp_size, cp_rank, s, b, h, d, d2, qkv_split_arg_list[0], qkv_split_arg_list[1], qkv_split_arg_list[2], - at::cuda::getCurrentCUDAStream()); + getCurrentCUDAStream()); return qkv_grad_input; } diff --git a/transformer_engine/pytorch/csrc/extensions/attention.cpp b/transformer_engine/pytorch/csrc/extensions/attention.cpp index 7e8018b3fd..ade2ad80dc 100644 --- a/transformer_engine/pytorch/csrc/extensions/attention.cpp +++ b/transformer_engine/pytorch/csrc/extensions/attention.cpp @@ -8,12 +8,16 @@ #include "common.h" #include "pybind.h" +// This TU has helpers in the anonymous (global) namespace; bring the facade +// aliases/functions (Tensor, getCurrentCUDAStream, ...) into scope for them. +using namespace transformer_engine::pytorch; // NOLINT(build/namespaces) + namespace { constexpr int block_size = 512; // fast zero-fills of tensors -void mha_fill(const transformer_engine::TensorWrapper &self, const at::Tensor &start_index) { +void mha_fill(const transformer_engine::TensorWrapper &self, const Tensor &start_index) { std::vector shape = transformer_engine::pytorch::convertShape(self.shape()); auto max_tokens = shape[0]; @@ -32,7 +36,7 @@ void mha_fill(const transformer_engine::TensorWrapper &self, const at::Tensor &s size_t total_bytes = num_rows_to_zero * fcd_size * element_size_bits / 8; NVTE_SCOPED_GIL_RELEASE( - { nvte_memset(base_ptr, 0, total_bytes, at::cuda::getCurrentCUDAStream()); }); + { nvte_memset(base_ptr, 0, total_bytes, getCurrentCUDAStream()); }); } } // namespace @@ -55,13 +59,13 @@ NVTE_Fused_Attn_Backend get_fused_attn_backend( } // helper function for S and dP quantizers -std::tuple> quantizer_helper( +std::tuple> quantizer_helper( py::handle quantizer, const std::vector &shape, DType dtype, bool create_hp_tensor, - std::optional data) { + std::optional data) { std::unique_ptr T_quantizer = convert_quantizer(quantizer); TensorWrapper te_T; py::object py_T; - std::optional amax_buf; + std::optional amax_buf; if (quantizer.is_none()) { // high precision auto *none_quantizer = dynamic_cast(T_quantizer.get()); @@ -116,18 +120,18 @@ std::vector fused_attn_fwd( bool set_zero, NVTE_QKV_Layout qkv_layout, NVTE_QKV_Format o_format, NVTE_QKV_Format qkv_scale_inv_format, NVTE_Bias_Type bias_type, NVTE_Mask_Type attn_mask_type, NVTE_Softmax_Type softmax_type, const std::vector window_size, - bool bottom_right_diagonal, const at::Tensor cu_seqlens_q, const at::Tensor cu_seqlens_kv, - const py::handle Q, const py::handle K, const py::handle V, const at::ScalarType fake_dtype, - const std::optional cu_seqlens_q_padded, - const std::optional cu_seqlens_kv_padded, - const std::optional page_table_k, const std::optional page_table_v, - py::handle s_quantizer, py::handle o_quantizer, const std::optional Bias, - const std::optional SoftmaxOffset, const std::optional rng_gen, + bool bottom_right_diagonal, const Tensor cu_seqlens_q, const Tensor cu_seqlens_kv, + const py::handle Q, const py::handle K, const py::handle V, const ScalarType fake_dtype, + const std::optional cu_seqlens_q_padded, + const std::optional cu_seqlens_kv_padded, + const std::optional page_table_k, const std::optional page_table_v, + py::handle s_quantizer, py::handle o_quantizer, const std::optional Bias, + const std::optional SoftmaxOffset, const std::optional rng_gen, size_t rng_elts_per_thread, bool return_max_logit, bool cuda_graph) { // Ensure that cuDNN handle is created on the correct device, // overriding torch.cuda.set_device calls from user side. // Assumes all tensors passed are on the same device. - at::cuda::CUDAGuard device_guard(cu_seqlens_q.device()); + CUDAGuard device_guard(cu_seqlens_q.device()); auto none = py::none(); @@ -165,14 +169,14 @@ std::vector fused_attn_fwd( // FP8 if (set_zero && (o_format == NVTE_QKV_Format::NVTE_THD)) { if ((h * d) % block_size == 0) { - mha_fill(te_O, cu_seqlens_q.index({torch::indexing::Slice(-1, torch::indexing::None)})); + mha_fill(te_O, cu_seqlens_q.index({indexing::Slice(-1, indexing::None)})); } else { - te_O.zero_(at::cuda::getCurrentCUDAStream()); + te_O.zero_(getCurrentCUDAStream()); } } } else if (qkv_type == DType::kBFloat16 || qkv_type == DType::kFloat16) { if (o_format == NVTE_QKV_Format::NVTE_THD) { - te_O.zero_(at::cuda::getCurrentCUDAStream()); + te_O.zero_(getCurrentCUDAStream()); } } else { NVTE_ERROR("Fused attention only supports FP8 and BF16/FP16 data types. \n"); @@ -228,11 +232,11 @@ std::vector fused_attn_fwd( } // extract rng seed and offset - auto gen = at::get_generator_or_default( - rng_gen, at::cuda::detail::getDefaultCUDAGenerator()); - at::PhiloxCudaState philox_args = init_philox_state(gen, rng_elts_per_thread); - auto options = torch::TensorOptions().dtype(torch::kInt64).device(torch::kCUDA); - auto rng_state = torch::empty({2}, options); + auto gen = get_generator_or_default( + rng_gen, getDefaultCUDAGenerator()); + PhiloxCudaState philox_args = init_philox_state(gen, rng_elts_per_thread); + auto options = TensorOptions().dtype(kInt64).device(kCUDA); + auto rng_state = empty({2}, options); philox_unpack(philox_args, static_cast(rng_state.data_ptr())); auto te_rng_state = makeTransformerEngineTensor(rng_state); @@ -252,7 +256,7 @@ std::vector fused_attn_fwd( te_page_table_v.data(), te_rng_state.data(), max_seqlen_q, max_seqlen_kv, is_training, return_max_logit, cuda_graph, attn_scale, p_dropout, qkv_layout, o_format, qkv_scale_inv_format, bias_type, attn_mask_type, softmax_type, window_size[0], - window_size[1], bottom_right_diagonal, workspace.data(), at::cuda::getCurrentCUDAStream()); + window_size[1], bottom_right_diagonal, workspace.data(), getCurrentCUDAStream()); }); // allocate memory for workspace and auxiliary output tensors @@ -263,7 +267,7 @@ std::vector fused_attn_fwd( // output_tensors = [O, nvte_aux_tensor_pack.tensors] std::vector output_tensors; output_tensors.push_back(py_O); - auto set_tensor_param = [&](size_t i, const at::Tensor &output_tensor) { + auto set_tensor_param = [&](size_t i, const Tensor &output_tensor) { output_tensors.push_back(py::cast(output_tensor)); NVTEBasicTensor temp_data = {output_tensor.data_ptr(), nvte_tensor_type(nvte_aux_tensor_pack.tensors[i]), @@ -274,7 +278,7 @@ std::vector fused_attn_fwd( // f16_arbitrary: S [b, h, sq, 1]/[tq, h, 1], (optional) Max [b, h, sq, 1]/[tq, h, 1], rng_state [2], (optional) Bias [1, h, sq, skv], (optional) SoftmaxOffset [1, h, 1, 1] // fp8 : S [b, h, sq, 1], rng_state [2] size_t i = 0; - at::Tensor output_tensor; + Tensor output_tensor; // intermediate softmax stats tensor S output_tensor = allocateSpace(nvte_shape_to_vector(nvte_tensor_shape(nvte_aux_tensor_pack.tensors[i])), @@ -309,7 +313,7 @@ std::vector fused_attn_fwd( te_page_table_v.data(), te_rng_state.data(), max_seqlen_q, max_seqlen_kv, is_training, return_max_logit, cuda_graph, attn_scale, p_dropout, qkv_layout, o_format, qkv_scale_inv_format, bias_type, attn_mask_type, softmax_type, window_size[0], - window_size[1], bottom_right_diagonal, workspace.data(), at::cuda::getCurrentCUDAStream()); + window_size[1], bottom_right_diagonal, workspace.data(), getCurrentCUDAStream()); }); // destroy tensor wrappers, but not allocated memory @@ -326,12 +330,12 @@ std::vector fused_attn_bwd( NVTE_QKV_Layout dqkv_layout, NVTE_QKV_Format qkv_scale_inv_format, NVTE_QKV_Format do_scale_inv_format, NVTE_Bias_Type bias_type, NVTE_Mask_Type attn_mask_type, NVTE_Softmax_Type softmax_type, const std::vector window_size, - bool bottom_right_diagonal, bool deterministic, const at::Tensor cu_seqlens_q, - const at::Tensor cu_seqlens_kv, const py::handle Q, const py::handle K, const py::handle V, - const py::handle O, const py::handle dO, const at::ScalarType fake_dtype, - const std::vector Aux_CTX_Tensors, - const std::optional cu_seqlens_q_padded, - const std::optional cu_seqlens_kv_padded, py::handle s_quantizer, + bool bottom_right_diagonal, bool deterministic, const Tensor cu_seqlens_q, + const Tensor cu_seqlens_kv, const py::handle Q, const py::handle K, const py::handle V, + const py::handle O, const py::handle dO, const ScalarType fake_dtype, + const std::vector Aux_CTX_Tensors, + const std::optional cu_seqlens_q_padded, + const std::optional cu_seqlens_kv_padded, py::handle s_quantizer, py::handle dp_quantizer, py::handle dqkv_quantizer, bool cuda_graph) { auto none = py::none(); @@ -351,7 +355,7 @@ std::vector fused_attn_bwd( // create dQ, dK, dV tensors TensorWrapper te_dQ, te_dK, te_dV; py::object py_dQ, py_dK, py_dV; - std::optional dq_amax_buf, dk_amax_buf, dv_amax_buf; + std::optional dq_amax_buf, dk_amax_buf, dv_amax_buf; std::unique_ptr dQKV_quantizer = convert_quantizer(dqkv_quantizer); std::vector q_shape = convertShape(te_Q.shape()); std::vector k_shape = convertShape(te_K.shape()); @@ -373,13 +377,13 @@ std::vector fused_attn_bwd( AttentionShape v_parsed(kv_format, v_shape.data()); size_t d_v = v_parsed.d(); v_parsed.to_format(dkv_format, dV_shape.data()); - at::Tensor dQ, dK, dV, dQKV, dKV; + Tensor dQ, dK, dV, dQKV, dKV; // FP16/BF16: dqkv_fake_dtype = kFloat16/kBFloat16, dQ/dK/dV.dtype = torch.float16/torch.bfloat16 // FP8DS: dqkv_fake_dtype = kFloat16/kBFloat16, dQ/dK/dV.dtype = torch.uint8 // FP8CS/MXFP8: dqkv_fake_dtype = kFloat16/kBFloat16, dQ/dK/dV.dtype = torch.float16/torch.bfloat16 - auto options = torch::TensorOptions().dtype(fake_dtype).device(torch::kCUDA); + auto options = TensorOptions().dtype(fake_dtype).device(kCUDA); if (detail::IsFloat8Quantizers(dqkv_quantizer.ptr())) { - options = options.dtype(torch::kUInt8); + options = options.dtype(kUInt8); } NVTE_QKV_Layout_Group layout_group = nvte_get_qkv_layout_group(dqkv_layout); @@ -388,70 +392,70 @@ std::vector fused_attn_bwd( case NVTE_QKV_Layout_Group::NVTE_3HD: tmp_shape = std::vector{dQ_shape.begin(), dQ_shape.end()}; tmp_shape.insert(tmp_shape.begin() + tmp_shape.size() - 2, int64_t(3)); - dQKV = torch::empty(c10::IntArrayRef(tmp_shape), options); - dQ = dQKV.index({"...", torch::indexing::Slice(0, 1, 1), - torch::indexing::Slice(0, torch::indexing::None, 1), - torch::indexing::Slice(0, torch::indexing::None, 1)}) + dQKV = empty(IntArrayRef(tmp_shape), options); + dQ = dQKV.index({"...", indexing::Slice(0, 1, 1), + indexing::Slice(0, indexing::None, 1), + indexing::Slice(0, indexing::None, 1)}) .squeeze(tmp_shape.size() - 3); - dK = dQKV.index({"...", torch::indexing::Slice(1, 2, 1), - torch::indexing::Slice(0, torch::indexing::None, 1), - torch::indexing::Slice(0, torch::indexing::None, 1)}) + dK = dQKV.index({"...", indexing::Slice(1, 2, 1), + indexing::Slice(0, indexing::None, 1), + indexing::Slice(0, indexing::None, 1)}) .squeeze(tmp_shape.size() - 3); - dV = dQKV.index({"...", torch::indexing::Slice(2, torch::indexing::None, 1), - torch::indexing::Slice(0, torch::indexing::None, 1), - torch::indexing::Slice(0, torch::indexing::None, 1)}) + dV = dQKV.index({"...", indexing::Slice(2, indexing::None, 1), + indexing::Slice(0, indexing::None, 1), + indexing::Slice(0, indexing::None, 1)}) .squeeze(tmp_shape.size() - 3); break; case NVTE_QKV_Layout_Group::NVTE_H3D: tmp_shape = std::vector{dQ_shape.begin(), dQ_shape.end()}; tmp_shape.insert(tmp_shape.begin() + tmp_shape.size() - 1, int64_t(3)); - dQKV = torch::empty(c10::IntArrayRef(tmp_shape), options); - dQ = dQKV.index({"...", torch::indexing::Slice(0, 1, 1), - torch::indexing::Slice(0, torch::indexing::None, 1)}) + dQKV = empty(IntArrayRef(tmp_shape), options); + dQ = dQKV.index({"...", indexing::Slice(0, 1, 1), + indexing::Slice(0, indexing::None, 1)}) .squeeze(tmp_shape.size() - 2); - dK = dQKV.index({"...", torch::indexing::Slice(1, 2, 1), - torch::indexing::Slice(0, torch::indexing::None, 1)}) + dK = dQKV.index({"...", indexing::Slice(1, 2, 1), + indexing::Slice(0, indexing::None, 1)}) .squeeze(tmp_shape.size() - 2); - dV = dQKV.index({"...", torch::indexing::Slice(2, torch::indexing::None, 1), - torch::indexing::Slice(0, torch::indexing::None, 1)}) + dV = dQKV.index({"...", indexing::Slice(2, indexing::None, 1), + indexing::Slice(0, indexing::None, 1)}) .squeeze(tmp_shape.size() - 2); break; case NVTE_QKV_Layout_Group::NVTE_HD_2HD: tmp_shape = std::vector(dQ_shape.begin(), dQ_shape.end()); - dQ = torch::empty(tmp_shape, options); + dQ = empty(tmp_shape, options); tmp_shape = std::vector{dK_shape.begin(), dK_shape.end()}; tmp_shape.insert(tmp_shape.begin() + tmp_shape.size() - 2, int64_t(2)); - dKV = torch::empty(c10::IntArrayRef(tmp_shape), options); - dK = dKV.index({"...", torch::indexing::Slice(0, 1, 1), - torch::indexing::Slice(0, torch::indexing::None, 1), - torch::indexing::Slice(0, torch::indexing::None, 1)}) + dKV = empty(IntArrayRef(tmp_shape), options); + dK = dKV.index({"...", indexing::Slice(0, 1, 1), + indexing::Slice(0, indexing::None, 1), + indexing::Slice(0, indexing::None, 1)}) .squeeze(tmp_shape.size() - 3); - dV = dKV.index({"...", torch::indexing::Slice(1, torch::indexing::None, 1), - torch::indexing::Slice(0, torch::indexing::None, 1), - torch::indexing::Slice(0, torch::indexing::None, 1)}) + dV = dKV.index({"...", indexing::Slice(1, indexing::None, 1), + indexing::Slice(0, indexing::None, 1), + indexing::Slice(0, indexing::None, 1)}) .squeeze(tmp_shape.size() - 3); break; case NVTE_QKV_Layout_Group::NVTE_HD_H2D: tmp_shape = std::vector(dQ_shape.begin(), dQ_shape.end()); - dQ = torch::empty(tmp_shape, options); + dQ = empty(tmp_shape, options); tmp_shape = std::vector{dK_shape.begin(), dK_shape.end()}; tmp_shape.insert(tmp_shape.begin() + tmp_shape.size() - 1, int64_t(2)); - dKV = torch::empty(c10::IntArrayRef(tmp_shape), options); - dK = dKV.index({"...", torch::indexing::Slice(0, 1, 1), - torch::indexing::Slice(0, torch::indexing::None, 1)}) + dKV = empty(IntArrayRef(tmp_shape), options); + dK = dKV.index({"...", indexing::Slice(0, 1, 1), + indexing::Slice(0, indexing::None, 1)}) .squeeze(tmp_shape.size() - 2); - dV = dKV.index({"...", torch::indexing::Slice(1, torch::indexing::None, 1), - torch::indexing::Slice(0, torch::indexing::None, 1)}) + dV = dKV.index({"...", indexing::Slice(1, indexing::None, 1), + indexing::Slice(0, indexing::None, 1)}) .squeeze(tmp_shape.size() - 2); break; case NVTE_QKV_Layout_Group::NVTE_HD_HD_HD: case NVTE_QKV_Layout_Group::NVTE_SD_SD_SD: tmp_shape = std::vector(dQ_shape.begin(), dQ_shape.end()); - dQ = torch::empty(tmp_shape, options); + dQ = empty(tmp_shape, options); tmp_shape = std::vector(dK_shape.begin(), dK_shape.end()); - dK = torch::empty(tmp_shape, options); + dK = empty(tmp_shape, options); tmp_shape = std::vector(dV_shape.begin(), dV_shape.end()); - dV = torch::empty(tmp_shape, options); + dV = empty(tmp_shape, options); break; default: NVTE_ERROR("QKV layout not supported!"); @@ -470,7 +474,7 @@ std::vector fused_attn_bwd( if (set_zero) { if (dq_format == NVTE_QKV_Format::NVTE_THD) { if (((h_q * d_qk) % block_size == 0) && dQ.is_contiguous()) { - mha_fill(te_dQ, cu_seqlens_q.index({torch::indexing::Slice(-1, torch::indexing::None)})); + mha_fill(te_dQ, cu_seqlens_q.index({indexing::Slice(-1, indexing::None)})); } else { dQ.fill_(0); } @@ -478,8 +482,8 @@ std::vector fused_attn_bwd( if (dkv_format == NVTE_QKV_Format::NVTE_THD) { if (((h_kv * d_qk) % block_size == 0) && ((h_kv * d_v) % block_size == 0) && dK.is_contiguous() && dV.is_contiguous()) { - mha_fill(te_dK, cu_seqlens_kv.index({torch::indexing::Slice(-1, torch::indexing::None)})); - mha_fill(te_dV, cu_seqlens_kv.index({torch::indexing::Slice(-1, torch::indexing::None)})); + mha_fill(te_dK, cu_seqlens_kv.index({indexing::Slice(-1, indexing::None)})); + mha_fill(te_dV, cu_seqlens_kv.index({indexing::Slice(-1, indexing::None)})); } else { dK.fill_(0); dV.fill_(0); @@ -541,15 +545,15 @@ std::vector fused_attn_bwd( } // create dBias the same shape as Bias - at::Tensor dBias; + Tensor dBias; TensorWrapper te_dBias; if ((bias_type != NVTE_NO_BIAS) && (bias_type != NVTE_ALIBI)) { if (nvte_aux_tensor_pack.size >= 2) { std::vector bias_shape(Aux_CTX_Tensors[nvte_aux_tensor_pack.size - 1].sizes().vec()); - dBias = torch::empty(bias_shape, options); + dBias = empty(bias_shape, options); te_dBias = makeTransformerEngineTensor(dBias); } else { - dBias = torch::empty({1, static_cast(h_q), static_cast(max_seqlen_q), + dBias = empty({1, static_cast(h_q), static_cast(max_seqlen_q), static_cast(max_seqlen_kv)}, options); te_dBias = makeTransformerEngineTensor(dBias); @@ -560,11 +564,11 @@ std::vector fused_attn_bwd( } // create dSoftmaxOffset in the same shape as SoftmaxOffset - at::Tensor dSoftmaxOffset; + Tensor dSoftmaxOffset; TensorWrapper te_dSoftmaxOffset; if (softmax_type != NVTE_VANILLA_SOFTMAX) { - options = torch::TensorOptions().dtype(at::kFloat).device(torch::kCUDA); - dSoftmaxOffset = torch::empty({1, static_cast(h_q), 1, 1}, options); + options = TensorOptions().dtype(kFloat).device(kCUDA); + dSoftmaxOffset = empty({1, static_cast(h_q), 1, 1}, options); te_dSoftmaxOffset = makeTransformerEngineTensor(dSoftmaxOffset); } @@ -581,7 +585,7 @@ std::vector fused_attn_bwd( attn_scale, p_dropout, qkv_layout, o_format, do_format, dqkv_layout, qkv_scale_inv_format, do_scale_inv_format, bias_type, attn_mask_type, softmax_type, window_size[0], window_size[1], bottom_right_diagonal, deterministic, cuda_graph, workspace.data(), - at::cuda::getCurrentCUDAStream()); + getCurrentCUDAStream()); }); // allocate memory for workspace @@ -599,7 +603,7 @@ std::vector fused_attn_bwd( attn_scale, p_dropout, qkv_layout, o_format, do_format, dqkv_layout, qkv_scale_inv_format, do_scale_inv_format, bias_type, attn_mask_type, softmax_type, window_size[0], window_size[1], bottom_right_diagonal, deterministic, cuda_graph, workspace.data(), - at::cuda::getCurrentCUDAStream()); + getCurrentCUDAStream()); }); // destroy tensor wrappers @@ -608,10 +612,10 @@ std::vector fused_attn_bwd( return {py_dQ, py_dK, py_dV, py::cast(dBias), py::cast(dSoftmaxOffset)}; } -at::Tensor fa_prepare_fwd(at::Tensor qkvi) { +Tensor fa_prepare_fwd(Tensor qkvi) { NVTE_CHECK(qkvi.dim() == 4, "Expected 4-dim tensor."); - NVTE_CHECK(qkvi.scalar_type() == at::ScalarType::Half || - qkvi.scalar_type() == at::ScalarType::BFloat16); + NVTE_CHECK(qkvi.scalar_type() == ScalarType::Half || + qkvi.scalar_type() == ScalarType::BFloat16); NVTE_CHECK(qkvi.stride(3) == 1, "Wrong stride."); NVTE_CHECK(qkvi.stride(2) == 3 * qkvi.size(3), "Wrong stride."); NVTE_CHECK(qkvi.stride(1) == 3 * qkvi.size(3) * qkvi.size(2), "Wrong stride."); @@ -619,31 +623,31 @@ at::Tensor fa_prepare_fwd(at::Tensor qkvi) { // [s, b, n, h * 3] -> [3, b, s, n, h] std::vector shape = {3, qkvi.size(1), qkvi.size(0), qkvi.size(2), qkvi.size(3)}; - at::Tensor qkv = at::empty(shape, at::CUDA(qkvi.scalar_type())); + Tensor qkv = empty(shape, CUDA(qkvi.scalar_type())); auto te_qkvi = makeTransformerEngineTensor(qkvi); auto te_qkv = makeTransformerEngineTensor(qkv); - nvte_prepare_flash_attn_fwd(te_qkvi.data(), te_qkv.data(), at::cuda::getCurrentCUDAStream()); + nvte_prepare_flash_attn_fwd(te_qkvi.data(), te_qkv.data(), getCurrentCUDAStream()); return qkv; } -at::Tensor fa_prepare_bwd(at::Tensor q, at::Tensor k, at::Tensor v) { +Tensor fa_prepare_bwd(Tensor q, Tensor k, Tensor v) { NVTE_CHECK(q.is_contiguous()); NVTE_CHECK(k.is_contiguous()); NVTE_CHECK(v.is_contiguous()); NVTE_CHECK(q.dim() == 4, "Expected 4-dim tensor."); NVTE_CHECK(k.dim() == 4, "Expected 4-dim tensor."); NVTE_CHECK(v.dim() == 4, "Expected 4-dim tensor."); - NVTE_CHECK(q.scalar_type() == at::ScalarType::Half || - q.scalar_type() == at::ScalarType::BFloat16); + NVTE_CHECK(q.scalar_type() == ScalarType::Half || + q.scalar_type() == ScalarType::BFloat16); NVTE_CHECK(k.scalar_type() == q.scalar_type()); NVTE_CHECK(v.scalar_type() == q.scalar_type()); // 3 x [s, b, n, h] -> [b, s, n, 3 * h] std::vector shape = {q.size(1), q.size(0), q.size(2), 3 * q.size(3)}; - at::Tensor qkv = at::empty(shape, at::CUDA(q.scalar_type())); + Tensor qkv = empty(shape, CUDA(q.scalar_type())); auto te_q = makeTransformerEngineTensor(q); auto te_k = makeTransformerEngineTensor(k); @@ -651,14 +655,14 @@ at::Tensor fa_prepare_bwd(at::Tensor q, at::Tensor k, at::Tensor v) { auto te_qkv = makeTransformerEngineTensor(qkv); nvte_prepare_flash_attn_bwd(te_q.data(), te_k.data(), te_v.data(), te_qkv.data(), - at::cuda::getCurrentCUDAStream()); + getCurrentCUDAStream()); return qkv; } -std::vector> multi_tensor_transpose_to_bhsd( - std::vector> inputs, const std::string &original_format, - std::vector> outputs) { +std::vector> multi_tensor_transpose_to_bhsd( + std::vector> inputs, const std::string &original_format, + std::vector> outputs) { NVTE_CHECK(original_format == "sbhd" || original_format == "bshd", "multi_tensor_transpose_to_bhsd: only BSHD/SBHD -> BHSD is currently supported. " "Got original_format=\"", @@ -674,7 +678,7 @@ std::vector> multi_tensor_transpose_to_bhsd( } std::vector te_ins, te_outs; - std::vector> result(inputs.size(), std::nullopt); + std::vector> result(inputs.size(), std::nullopt); for (size_t i = 0; i < inputs.size(); ++i) { if (!inputs[i].has_value()) continue; @@ -683,12 +687,12 @@ std::vector> multi_tensor_transpose_to_bhsd( NVTE_CHECK(input.is_cuda() && input.dim() == 4, "multi_tensor_transpose_to_bhsd: input ", i, " must be a 4D CUDA tensor."); input = input.contiguous(); - NVTE_CHECK(input.scalar_type() == at::ScalarType::Half || - input.scalar_type() == at::ScalarType::BFloat16 || - input.scalar_type() == at::ScalarType::Byte, + NVTE_CHECK(input.scalar_type() == ScalarType::Half || + input.scalar_type() == ScalarType::BFloat16 || + input.scalar_type() == ScalarType::Byte, "multi_tensor_transpose_to_bhsd: unsupported dtype at index ", i, "."); - at::Tensor output; + Tensor output; if (has_outputs && outputs[i].has_value()) { output = outputs[i].value(); } else { @@ -704,7 +708,7 @@ std::vector> multi_tensor_transpose_to_bhsd( H = input.size(2); D = input.size(3); } - output = at::empty({B, H, S, D}, input.options()); + output = empty({B, H, S, D}, input.options()); } te_ins.push_back(makeTransformerEngineTensor(input)); @@ -719,20 +723,20 @@ std::vector> multi_tensor_transpose_to_bhsd( nvte_outs[j] = te_outs[j].data(); } nvte_multi_tensor_transpose_to_bhsd(nvte_ins.data(), nvte_outs.data(), te_ins.size(), - original_format_enum, at::cuda::getCurrentCUDAStream()); + original_format_enum, getCurrentCUDAStream()); } return result; } -std::vector multi_tensor_pad_last_dim(std::vector inputs, +std::vector multi_tensor_pad_last_dim(std::vector inputs, int64_t alignment) { const auto align = static_cast(alignment); NVTE_CHECK(align > 0, "multi_tensor_pad_last_dim: alignment must be > 0."); NVTE_CHECK(!inputs.empty(), "multi_tensor_pad_last_dim: inputs must not be empty."); - auto stream = at::cuda::getCurrentCUDAStream(); - std::vector outputs; + auto stream = getCurrentCUDAStream(); + std::vector outputs; outputs.reserve(inputs.size()); std::vector kernel_indices; @@ -756,7 +760,7 @@ std::vector multi_tensor_pad_last_dim(std::vector inputs continue; } - at::Tensor output = at::empty({rows, padded_cols}, input.options()); + Tensor output = empty({rows, padded_cols}, input.options()); outputs.push_back(output); kernel_indices.push_back(outputs.size() - 1); } @@ -789,10 +793,10 @@ std::vector multi_tensor_pad_last_dim(std::vector inputs * Support THD format for Context Parallel: Read the half of a THD tensor **************************************************************************************************/ -at::Tensor thd_read_half_tensor(const at::Tensor &tensor, const at::Tensor &cu_seqlens, +Tensor thd_read_half_tensor(const Tensor &tensor, const Tensor &cu_seqlens, int half_idx) { NVTE_CHECK(tensor.dim() == 3 || tensor.dim() == 4); - NVTE_CHECK(cu_seqlens.scalar_type() == at::ScalarType::Int); + NVTE_CHECK(cu_seqlens.scalar_type() == ScalarType::Int); NVTE_CHECK(cu_seqlens.dim() == 1); NVTE_CHECK(cu_seqlens.size(0) >= 2); @@ -802,7 +806,7 @@ at::Tensor thd_read_half_tensor(const at::Tensor &tensor, const at::Tensor &cu_s int num_heads = tensor.size(seq_dim + 1); int dim_per_head = tensor.size(seq_dim + 2); - int hidden_size_in_bytes = num_heads * dim_per_head * c10::elementSize(tensor.scalar_type()); + int hidden_size_in_bytes = num_heads * dim_per_head * elementSize(tensor.scalar_type()); // For 128-bits load/store NVTE_CHECK(hidden_size_in_bytes % 16 == 0); @@ -813,14 +817,14 @@ at::Tensor thd_read_half_tensor(const at::Tensor &tensor, const at::Tensor &cu_s shape[i] = tensor.size(i); } shape[seq_dim] /= 2; - at::Tensor half = at::empty(shape, at::CUDA(tensor.scalar_type())); + Tensor half = empty(shape, CUDA(tensor.scalar_type())); auto te_tensor = makeTransformerEngineTensor(tensor); auto te_cu_seqlens = makeTransformerEngineTensor(cu_seqlens); auto te_half = makeTransformerEngineTensor(half); nvte_cp_thd_read_half_tensor(te_tensor.data(), te_cu_seqlens.data(), te_half.data(), half_idx, - at::cuda::getCurrentCUDAStream()); + getCurrentCUDAStream()); return half; } @@ -829,11 +833,11 @@ at::Tensor thd_read_half_tensor(const at::Tensor &tensor, const at::Tensor &cu_s * Support THD format for Context Parallel: softmax_lse related operations **************************************************************************************************/ -void thd_second_half_lse_correction(at::Tensor lse, const at::Tensor &lse_per_step, - const at::Tensor &cu_seqlens, bool lse_packed) { - NVTE_CHECK(lse.scalar_type() == at::ScalarType::Float); - NVTE_CHECK(lse_per_step.scalar_type() == at::ScalarType::Float); - NVTE_CHECK(cu_seqlens.scalar_type() == at::ScalarType::Int); +void thd_second_half_lse_correction(Tensor lse, const Tensor &lse_per_step, + const Tensor &cu_seqlens, bool lse_packed) { + NVTE_CHECK(lse.scalar_type() == ScalarType::Float); + NVTE_CHECK(lse_per_step.scalar_type() == ScalarType::Float); + NVTE_CHECK(cu_seqlens.scalar_type() == ScalarType::Int); NVTE_CHECK(cu_seqlens.dim() == 1); int batch, num_heads, lse_seqlen, second_half_lse_seqlen; @@ -870,13 +874,13 @@ void thd_second_half_lse_correction(at::Tensor lse, const at::Tensor &lse_per_st nvte_cp_thd_second_half_lse_correction(te_lse.data(), te_lse_per_step.data(), te_cu_seqlens.data(), lse_packed, - at::cuda::getCurrentCUDAStream()); + getCurrentCUDAStream()); } -at::Tensor thd_read_second_half_lse(const at::Tensor &lse, const at::Tensor &cu_seqlens, +Tensor thd_read_second_half_lse(const Tensor &lse, const Tensor &cu_seqlens, bool lse_packed, int second_half_lse_seqlen) { - NVTE_CHECK(lse.scalar_type() == at::ScalarType::Float); - NVTE_CHECK(cu_seqlens.scalar_type() == at::ScalarType::Int); + NVTE_CHECK(lse.scalar_type() == ScalarType::Float); + NVTE_CHECK(cu_seqlens.scalar_type() == ScalarType::Int); NVTE_CHECK(cu_seqlens.dim() == 1); int batch, num_heads, lse_seqlen; @@ -905,7 +909,7 @@ at::Tensor thd_read_second_half_lse(const at::Tensor &lse, const at::Tensor &cu_ shape = {batch, num_heads, second_half_lse_seqlen}; } - at::Tensor half_lse = at::zeros(shape, at::CUDA(lse.scalar_type())); + Tensor half_lse = zeros(shape, CUDA(lse.scalar_type())); auto te_lse = makeTransformerEngineTensor(lse); auto te_cu_seqlens = makeTransformerEngineTensor(cu_seqlens); @@ -913,7 +917,7 @@ at::Tensor thd_read_second_half_lse(const at::Tensor &lse, const at::Tensor &cu_ nvte_cp_thd_read_second_half_lse(te_lse.data(), te_cu_seqlens.data(), te_half_lse.data(), lse_packed, second_half_lse_seqlen, - at::cuda::getCurrentCUDAStream()); + getCurrentCUDAStream()); return half_lse; } @@ -922,8 +926,8 @@ at::Tensor thd_read_second_half_lse(const at::Tensor &lse, const at::Tensor &cu_ * Support THD format for Context Parallel: Out correction in forward **************************************************************************************************/ -void thd_out_correction(at::Tensor out, const at::Tensor &out_per_step, const at::Tensor &lse, - const at::Tensor &lse_per_step, const at::Tensor &cu_seqlens, +void thd_out_correction(Tensor out, const Tensor &out_per_step, const Tensor &lse, + const Tensor &lse_per_step, const Tensor &cu_seqlens, bool only_second_half, bool lse_packed) { auto te_out = makeTransformerEngineTensor(out); auto te_out_per_step = makeTransformerEngineTensor(out_per_step); @@ -932,31 +936,31 @@ void thd_out_correction(at::Tensor out, const at::Tensor &out_per_step, const at auto te_cu_seqlens = makeTransformerEngineTensor(cu_seqlens); nvte_cp_thd_out_correction(te_out.data(), te_out_per_step.data(), te_lse.data(), te_lse_per_step.data(), te_cu_seqlens.data(), only_second_half, - lse_packed, at::cuda::getCurrentCUDAStream()); + lse_packed, getCurrentCUDAStream()); } /*************************************************************************************************** * Support THD format for Context Parallel: Gradients correction in backward **************************************************************************************************/ -void thd_grad_correction(at::Tensor grad, const at::Tensor &grad_per_step, - const at::Tensor &cu_seqlens, const std::string &first_half, +void thd_grad_correction(Tensor grad, const Tensor &grad_per_step, + const Tensor &cu_seqlens, const std::string &first_half, const std::string &second_half) { auto te_grad = makeTransformerEngineTensor(grad); auto te_grad_per_step = makeTransformerEngineTensor(grad_per_step); auto te_cu_seqlens = makeTransformerEngineTensor(cu_seqlens); nvte_cp_thd_grad_correction(te_grad.data(), te_grad_per_step.data(), te_cu_seqlens.data(), first_half.data(), second_half.data(), - at::cuda::getCurrentCUDAStream()); + getCurrentCUDAStream()); } /*************************************************************************************************** * Support THD format for Context Parallel: Generate partitioned indices for input tokens **************************************************************************************************/ -at::Tensor thd_get_partitioned_indices(const at::Tensor &cu_seqlens, int total_tokens, +Tensor thd_get_partitioned_indices(const Tensor &cu_seqlens, int total_tokens, int world_size, int rank) { - NVTE_CHECK(cu_seqlens.scalar_type() == at::ScalarType::Int); + NVTE_CHECK(cu_seqlens.scalar_type() == ScalarType::Int); NVTE_CHECK(cu_seqlens.dim() == 1); NVTE_CHECK(cu_seqlens.size(0) >= 2); NVTE_CHECK(rank >= 0 && rank < world_size); @@ -964,13 +968,13 @@ at::Tensor thd_get_partitioned_indices(const at::Tensor &cu_seqlens, int total_t NVTE_CHECK(total_tokens > 0 && total_tokens % (world_size * 2) == 0); std::vector shape = {total_tokens / world_size}; - at::Tensor output = at::empty(shape, at::CUDA(at::ScalarType::Int)); + Tensor output = empty(shape, CUDA(ScalarType::Int)); auto te_cu_seqlens = makeTransformerEngineTensor(cu_seqlens); auto te_output = makeTransformerEngineTensor(output); nvte_cp_thd_get_partitioned_indices(te_cu_seqlens.data(), te_output.data(), total_tokens, - world_size, rank, at::cuda::getCurrentCUDAStream()); + world_size, rank, getCurrentCUDAStream()); return output; } @@ -979,18 +983,18 @@ at::Tensor thd_get_partitioned_indices(const at::Tensor &cu_seqlens, int total_t * KV Cache: Convert a tensor from qkv_format = thd to qkv_format = bshd **************************************************************************************************/ -at::Tensor convert_thd_to_bshd(at::Tensor tensor, at::Tensor cu_seqlens, int b, int max_seq_len) { +Tensor convert_thd_to_bshd(Tensor tensor, Tensor cu_seqlens, int b, int max_seq_len) { int h = tensor.size(1); int d = tensor.size(2); std::vector shape = {b, max_seq_len, h, d}; - at::Tensor new_tensor = at::zeros(shape, at::CUDA(tensor.scalar_type())); + Tensor new_tensor = zeros(shape, CUDA(tensor.scalar_type())); auto te_tensor = makeTransformerEngineTensor(tensor); auto te_cu_seqlens = makeTransformerEngineTensor(cu_seqlens); auto te_new_tensor = makeTransformerEngineTensor(new_tensor); nvte_convert_thd_to_bshd(te_tensor.data(), te_cu_seqlens.data(), te_new_tensor.data(), b, - max_seq_len, at::cuda::getCurrentCUDAStream()); + max_seq_len, getCurrentCUDAStream()); return new_tensor; } @@ -999,24 +1003,24 @@ at::Tensor convert_thd_to_bshd(at::Tensor tensor, at::Tensor cu_seqlens, int b, * KV Cache: Convert a tensor from qkv_format = bshd to qkv_format = thd **************************************************************************************************/ -at::Tensor convert_bshd_to_thd(at::Tensor tensor, at::Tensor cu_seqlens, int t) { +Tensor convert_bshd_to_thd(Tensor tensor, Tensor cu_seqlens, int t) { int h = tensor.size(2); int d = tensor.size(3); std::vector shape = {t, h, d}; - at::Tensor new_tensor = at::zeros(shape, at::CUDA(tensor.scalar_type())); + Tensor new_tensor = zeros(shape, CUDA(tensor.scalar_type())); auto te_tensor = makeTransformerEngineTensor(tensor); auto te_cu_seqlens = makeTransformerEngineTensor(cu_seqlens); auto te_new_tensor = makeTransformerEngineTensor(new_tensor); nvte_convert_bshd_to_thd(te_tensor.data(), te_cu_seqlens.data(), te_new_tensor.data(), t, - at::cuda::getCurrentCUDAStream()); + getCurrentCUDAStream()); return new_tensor; } -void copy_to_kv_cache(at::Tensor new_k, at::Tensor new_v, at::Tensor k_cache, at::Tensor v_cache, - at::Tensor page_table, at::Tensor cu_new_lens, at::Tensor cu_cached_lens, +void copy_to_kv_cache(Tensor new_k, Tensor new_v, Tensor k_cache, Tensor v_cache, + Tensor page_table, Tensor cu_new_lens, Tensor cu_cached_lens, NVTE_QKV_Format qkv_format, int b, int max_ctx_len, int max_seq_len, int max_pages_per_seq, bool is_non_paged) { NVTE_CHECK(k_cache.scalar_type() == v_cache.scalar_type() && @@ -1038,7 +1042,7 @@ void copy_to_kv_cache(at::Tensor new_k, at::Tensor new_v, at::Tensor k_cache, at nvte_copy_to_kv_cache(te_new_k.data(), te_new_v.data(), te_k_cache.data(), te_v_cache.data(), te_page_table.data(), te_cu_new_lens.data(), te_cu_cached_lens.data(), qkv_format, b, max_ctx_len, max_seq_len, max_pages_per_seq, is_non_paged, - at::cuda::getCurrentCUDAStream()); + getCurrentCUDAStream()); } } // namespace transformer_engine::pytorch diff --git a/transformer_engine/pytorch/csrc/extensions/bias.cpp b/transformer_engine/pytorch/csrc/extensions/bias.cpp index 4a78dde388..83c63da08f 100644 --- a/transformer_engine/pytorch/csrc/extensions/bias.cpp +++ b/transformer_engine/pytorch/csrc/extensions/bias.cpp @@ -4,7 +4,6 @@ * See LICENSE for license information. ************************************************************************/ -#include #include #include @@ -19,7 +18,7 @@ namespace transformer_engine { namespace pytorch { -std::vector bgrad_quantize(const at::Tensor &grad_output, py::handle quantizer) { +std::vector bgrad_quantize(const Tensor &grad_output, py::handle quantizer) { using namespace transformer_engine::pytorch::detail; init_extension(); @@ -39,7 +38,7 @@ std::vector bgrad_quantize(const at::Tensor &grad_output, py::handle if (product(shape) == 0) { grad_bias_torch.zero_(); } else { - at::sum_out(grad_bias_torch, grad_output_torch.reshape({-1, bias_size}), {0}); + sum_out(grad_bias_torch, grad_output_torch.reshape({-1, bias_size}), {0}); } return {py::cast(std::move(grad_bias_torch)), py::cast(std::move(grad_output_torch))}; } @@ -57,7 +56,7 @@ std::vector bgrad_quantize(const at::Tensor &grad_output, py::handle // Check if fused kernel is supported bool with_fused_kernel = false; if (detail::IsFloat8Quantizers(quantizer.ptr())) { - auto prop = at::cuda::getCurrentDeviceProperties(); + auto prop = getCurrentDeviceProperties(); const size_t sm_arch = 10 * prop->major + prop->minor; if (sm_arch >= 100) { // Fused kernel for dbias + FP8 cast on SM arch 10.0+ @@ -73,15 +72,15 @@ std::vector bgrad_quantize(const at::Tensor &grad_output, py::handle // Apply unfused impl if fused kernel is not supported if (!with_fused_kernel) { - at::sum_out(grad_bias_torch, grad_output_torch.reshape({-1, bias_size}), {0}); + sum_out(grad_bias_torch, grad_output_torch.reshape({-1, bias_size}), {0}); quantizer_cpp->quantize(grad_output_nvte, grad_input_nvte); return {py::cast(std::move(grad_bias_torch)), std::move(grad_input_py)}; } // Query workspace size TensorWrapper workspace_nvte; - at::Tensor workspace_torch; - auto stream = at::cuda::getCurrentCUDAStream(); + Tensor workspace_torch; + auto stream = getCurrentCUDAStream(); NVTE_SCOPED_GIL_RELEASE({ nvte_quantize_dbias(grad_output_nvte.data(), grad_input_nvte.data(), grad_bias_nvte.data(), workspace_nvte.data(), stream); @@ -109,7 +108,7 @@ std::vector dact_dbias( void (*dact_dbias_func)(const NVTETensor, const NVTETensor, NVTETensor, NVTETensor, NVTETensor, cudaStream_t), void (*dact_func)(const NVTETensor, const NVTETensor, NVTETensor, cudaStream_t), - at::Tensor grad_output_torch, at::Tensor act_input_torch, py::handle quantizer_py) { + Tensor grad_output_torch, Tensor act_input_torch, py::handle quantizer_py) { using namespace transformer_engine::pytorch::detail; init_extension(); @@ -162,7 +161,7 @@ std::vector dact_dbias( } // Perform compute - auto stream = at::cuda::getCurrentCUDAStream(); + auto stream = getCurrentCUDAStream(); switch (impl) { case Impl::UNFUSED: // Unfused dact, dbias, quantize @@ -172,8 +171,8 @@ std::vector dact_dbias( NVTE_SCOPED_GIL_RELEASE({ dact_func(grad_output_nvte.data(), act_input_nvte.data(), temp_nvte.data(), stream); }); - const auto temp_torch = temp_py.cast(); - at::sum_out(grad_bias_torch, temp_torch.reshape({-1, bias_size}), {0}); + const auto temp_torch = temp_py.cast(); + sum_out(grad_bias_torch, temp_torch.reshape({-1, bias_size}), {0}); quantizer_cpp->quantize(temp_nvte, grad_input_nvte); break; } @@ -188,7 +187,7 @@ std::vector dact_dbias( }); // Allocate workspace - at::Tensor workspace_torch; + Tensor workspace_torch; if (workspace_nvte.ndim() > 0 && workspace_nvte.numel() > 0) { workspace_torch = allocateSpace(workspace_nvte.shape(), workspace_nvte.dtype()); workspace_nvte = makeTransformerEngineTensor( @@ -214,8 +213,8 @@ std::vector dact_dbias( NVTE_SCOPED_GIL_RELEASE({ dact_func(grad_output_nvte.data(), act_input_nvte.data(), temp_nvte.data(), stream); }); - const auto temp_torch = temp_py.cast(); - at::sum_out(grad_bias_torch, temp_torch.reshape({-1, bias_size}), {0}); + const auto temp_torch = temp_py.cast(); + sum_out(grad_bias_torch, temp_torch.reshape({-1, bias_size}), {0}); fp8_quantizer_cpp->quantize_with_amax(temp_nvte, grad_input_nvte, amax_buf); break; } @@ -231,8 +230,8 @@ std::vector dact_dbias( NVTE_SCOPED_GIL_RELEASE({ dact_func(grad_output_nvte.data(), act_input_nvte.data(), temp_nvte.data(), stream); }); - const auto temp_torch = temp_py.cast(); - at::sum_out(grad_bias_torch, temp_torch.reshape({-1, bias_size}), {0}); + const auto temp_torch = temp_py.cast(); + sum_out(grad_bias_torch, temp_torch.reshape({-1, bias_size}), {0}); nvfp4_quantizer_cpp->quantize_with_amax(temp_nvte, grad_input_nvte); break; } @@ -245,27 +244,27 @@ std::vector dact_dbias( } // namespace -std::vector dbias_dgelu(const at::Tensor &grad_output, const at::Tensor &act_input, +std::vector dbias_dgelu(const Tensor &grad_output, const Tensor &act_input, py::handle quantizer) { return dact_dbias(nvte_quantize_dbias_dgelu, nvte_dgelu, grad_output, act_input, quantizer); } -std::vector dbias_dsilu(const at::Tensor &grad_output, const at::Tensor &act_input, +std::vector dbias_dsilu(const Tensor &grad_output, const Tensor &act_input, py::handle quantizer) { return dact_dbias(nvte_quantize_dbias_dsilu, nvte_dsilu, grad_output, act_input, quantizer); } -std::vector dbias_drelu(const at::Tensor &grad_output, const at::Tensor &act_input, +std::vector dbias_drelu(const Tensor &grad_output, const Tensor &act_input, py::handle quantizer) { return dact_dbias(nvte_quantize_dbias_drelu, nvte_drelu, grad_output, act_input, quantizer); } -std::vector dbias_dqgelu(const at::Tensor &grad_output, const at::Tensor &act_input, +std::vector dbias_dqgelu(const Tensor &grad_output, const Tensor &act_input, py::handle quantizer) { return dact_dbias(nvte_quantize_dbias_dqgelu, nvte_dqgelu, grad_output, act_input, quantizer); } -std::vector dbias_dsrelu(const at::Tensor &grad_output, const at::Tensor &act_input, +std::vector dbias_dsrelu(const Tensor &grad_output, const Tensor &act_input, py::handle quantizer) { return dact_dbias(nvte_quantize_dbias_dsrelu, nvte_dsrelu, grad_output, act_input, quantizer); } diff --git a/transformer_engine/pytorch/csrc/extensions/cast.cpp b/transformer_engine/pytorch/csrc/extensions/cast.cpp index aab5a87b9a..898403d7f1 100644 --- a/transformer_engine/pytorch/csrc/extensions/cast.cpp +++ b/transformer_engine/pytorch/csrc/extensions/cast.cpp @@ -32,12 +32,12 @@ std::vector get_tensor_shape(const TensorWrapper &tensor) { } void allreduce_nvfp4_amax_tensors(NVFP4Quantizer *nvfp4_quantizer_cpp, - std::vector &&amax_tensors) { + std::vector &&amax_tensors) { if (!nvfp4_quantizer_cpp->with_amax_reduction || amax_tensors.empty()) { return; } - c10d::AllreduceCoalescedOptions opts; - opts.reduceOp = c10d::ReduceOp::MAX; + AllreduceCoalescedOptions opts; + opts.reduceOp = ReduceOp::MAX; NVTE_SCOPED_GIL_RELEASE({ nvfp4_quantizer_cpp->amax_reduction_group->allreduce_coalesced(amax_tensors, opts)->wait(); }); @@ -45,8 +45,8 @@ void allreduce_nvfp4_amax_tensors(NVFP4Quantizer *nvfp4_quantizer_cpp, } // namespace -py::object quantize(const at::Tensor &tensor, py::handle quantizer, const py::object &output, - std::optional noop_flag) { +py::object quantize(const Tensor &tensor, py::handle quantizer, const py::object &output, + std::optional noop_flag) { // Convert quantizer to C++ object auto quantizer_cpp = convert_quantizer(quantizer); @@ -83,9 +83,9 @@ py::object quantize(const at::Tensor &tensor, py::handle quantizer, const py::ob return output_py; } -py::object nvfp4_quantize_with_amax(const at::Tensor &tensor, py::handle quantizer, - const at::Tensor &rowwise_amax, - const at::Tensor &columnwise_amax) { +py::object nvfp4_quantize_with_amax(const Tensor &tensor, py::handle quantizer, + const Tensor &rowwise_amax, + const Tensor &columnwise_amax) { using namespace transformer_engine::pytorch::detail; init_extension(); @@ -93,7 +93,7 @@ py::object nvfp4_quantize_with_amax(const at::Tensor &tensor, py::handle quantiz NVTE_CHECK(rowwise_amax.is_cuda() && columnwise_amax.is_cuda(), "Precomputed amax tensors must be CUDA tensors."); NVTE_CHECK( - rowwise_amax.scalar_type() == at::kFloat && columnwise_amax.scalar_type() == at::kFloat, + rowwise_amax.scalar_type() == kFloat && columnwise_amax.scalar_type() == kFloat, "Precomputed amax tensors must be float32."); NVTE_CHECK(rowwise_amax.numel() == 1 && columnwise_amax.numel() == 1, "nvfp4_quantize_with_amax expects scalar rowwise and columnwise amaxes."); @@ -129,7 +129,7 @@ py::object nvfp4_quantize_with_amax(const at::Tensor &tensor, py::handle quantiz } py::object create_empty_quantized_tensor(py::handle quantizer, const std::vector &shape, - at::ScalarType dtype, at::Device device, bool pin_memory) { + ScalarType dtype, Device device, bool pin_memory) { auto quantizer_cpp = convert_quantizer(quantizer); auto te_dtype = GetTransformerEngineDType(dtype); auto [_, output_py] = quantizer_cpp->create_tensor(shape, te_dtype, device, pin_memory); @@ -159,18 +159,18 @@ void group_quantize_nvfp4_impl(const GroupedTensorWrapper &grouped_input_tensor, // stochastic rounding bool need_stochastic_rounding = nvfp4_quantizer_cpp->stochastic_rounding; - auto opts = at::TensorOptions().dtype(torch::kInt64).device(torch::kCUDA); - at::Tensor rng_states_tensor; // Declare tensor outside, do not allocate yet + auto opts = TensorOptions().dtype(kInt64).device(kCUDA); + Tensor rng_states_tensor; // Declare tensor outside, do not allocate yet TensorWrapper te_rng_state; if (need_stochastic_rounding) { // in fused kernel, one rng state will be used by the grouped kernel to generate random // number for different tensors in the group, so we only need to allocate one rng state const size_t rng_elts_per_thread = 1024 * num_tensors; - rng_states_tensor = torch::empty({2}, opts); - auto gen = at::get_generator_or_default( - std::nullopt, at::cuda::detail::getDefaultCUDAGenerator()); - at::PhiloxCudaState philox_args = init_philox_state(gen, rng_elts_per_thread); + rng_states_tensor = empty({2}, opts); + auto gen = get_generator_or_default( + std::nullopt, getDefaultCUDAGenerator()); + PhiloxCudaState philox_args = init_philox_state(gen, rng_elts_per_thread); philox_unpack(philox_args, static_cast(rng_states_tensor.data_ptr())); te_rng_state = makeTransformerEngineTensor(rng_states_tensor); @@ -193,7 +193,7 @@ void group_quantize_nvfp4_impl(const GroupedTensorWrapper &grouped_input_tensor, } // RHT cast fusion - auto tile_scheduler_workspace_torch = at::empty({1}, at::device(at::kCUDA).dtype(torch::kInt32)); + auto tile_scheduler_workspace_torch = empty({1}, TensorOptions().device(kCUDA).dtype(kInt32)); auto nvte_tile_scheduler_workspace = makeTransformerEngineTensor(tile_scheduler_workspace_torch); auto rht_matrix_nvte = makeTransformerEngineTensor(nvfp4_quantizer_cpp->rht_matrix); @@ -207,9 +207,9 @@ void group_quantize_nvfp4_impl(const GroupedTensorWrapper &grouped_input_tensor, } // namespace // NOTE: Only supports varying first dim. -py::object group_quantize(const at::Tensor &tensor, py::handle quantizer, const size_t num_tensors, - std::optional first_dims, - std::optional tensor_offsets) { +py::object group_quantize(const Tensor &tensor, py::handle quantizer, const size_t num_tensors, + std::optional first_dims, + std::optional tensor_offsets) { using namespace transformer_engine::pytorch::detail; init_extension(); @@ -263,14 +263,14 @@ py::object group_quantize(const at::Tensor &tensor, py::handle quantizer, const // NVFP4 grouped quantization NVFP4Quantizer *nvfp4_quantizer_cpp = static_cast(quantizer_cpp.get()); group_quantize_nvfp4_impl(grouped_input_tensor, grouped_output_tensor_cpp, - nvfp4_quantizer_cpp, at::cuda::getCurrentCUDAStream(), true); + nvfp4_quantizer_cpp, getCurrentCUDAStream(), true); break; } case GroupedQuantizationMode::MXFP8_GROUPED_QUANTIZE: { QuantizationConfigWrapper quant_config_cpp; NVTE_SCOPED_GIL_RELEASE({ nvte_group_quantize(grouped_input_tensor.data(), grouped_output_tensor_cpp.data(), - quant_config_cpp, at::cuda::getCurrentCUDAStream()); + quant_config_cpp, getCurrentCUDAStream()); }); break; } @@ -283,12 +283,12 @@ py::object group_quantize(const at::Tensor &tensor, py::handle quantizer, const return py::reinterpret_borrow(grouped_output_py); } -py::object nvfp4_group_quantize_with_amax(const at::Tensor &tensor, py::handle quantizer, +py::object nvfp4_group_quantize_with_amax(const Tensor &tensor, py::handle quantizer, const size_t num_tensors, - std::optional first_dims, - const at::Tensor &rowwise_amax, - const at::Tensor &columnwise_amax, - std::optional tensor_offsets) { + std::optional first_dims, + const Tensor &rowwise_amax, + const Tensor &columnwise_amax, + std::optional tensor_offsets) { using namespace transformer_engine::pytorch::detail; init_extension(); @@ -296,7 +296,7 @@ py::object nvfp4_group_quantize_with_amax(const at::Tensor &tensor, py::handle q NVTE_CHECK(rowwise_amax.is_cuda() && columnwise_amax.is_cuda(), "Precomputed amax tensors must be CUDA tensors."); NVTE_CHECK( - rowwise_amax.scalar_type() == at::kFloat && columnwise_amax.scalar_type() == at::kFloat, + rowwise_amax.scalar_type() == kFloat && columnwise_amax.scalar_type() == kFloat, "Precomputed amax tensors must be float32."); NVTE_CHECK(rowwise_amax.numel() == static_cast(num_tensors), "Rowwise amax must contain one value per group."); @@ -337,7 +337,7 @@ py::object nvfp4_group_quantize_with_amax(const at::Tensor &tensor, py::handle q grouped_output_py.attr("columnwise_amax") = py::cast(columnwise_amax); } - std::vector amax_tensors; + std::vector amax_tensors; if (grouped_output_tensor_cpp.get_amax().data_ptr != nullptr) { amax_tensors.push_back(rowwise_amax); } @@ -351,14 +351,14 @@ py::object nvfp4_group_quantize_with_amax(const at::Tensor &tensor, py::handle q } group_quantize_nvfp4_impl(grouped_input_tensor, grouped_output_tensor_cpp, nvfp4_quantizer_cpp, - at::cuda::getCurrentCUDAStream(), false); + getCurrentCUDAStream(), false); return py::reinterpret_borrow(grouped_output_py); } -py::object bgrad_group_quantize(const at::Tensor &tensor, py::handle quantizer, - const size_t num_tensors, std::optional first_dims, - std::optional tensor_offsets) { +py::object bgrad_group_quantize(const Tensor &tensor, py::handle quantizer, + const size_t num_tensors, std::optional first_dims, + std::optional tensor_offsets) { using namespace transformer_engine::pytorch::detail; init_extension(); @@ -388,8 +388,8 @@ py::object bgrad_group_quantize(const at::Tensor &tensor, py::handle quantizer, logical_last_dim); if (empty_input_buffer) { - at::Tensor dbias_torch = - at::zeros({static_cast(num_tensors), static_cast(logical_last_dim)}, + Tensor dbias_torch = + zeros({static_cast(num_tensors), static_cast(logical_last_dim)}, tensor.options()); return py::make_tuple(py::reinterpret_borrow(grouped_output_py), py::cast(std::move(dbias_torch))); @@ -397,20 +397,20 @@ py::object bgrad_group_quantize(const at::Tensor &tensor, py::handle quantizer, const std::vector dbias_logical_shape = {num_tensors, logical_last_dim}; GroupedTensorWrapper grouped_dbias(num_tensors, dbias_logical_shape, NVTE_DELAYED_TENSOR_SCALING); - at::Tensor dbias_torch = - at::empty({static_cast(num_tensors), static_cast(logical_last_dim)}, + Tensor dbias_torch = + empty({static_cast(num_tensors), static_cast(logical_last_dim)}, tensor.options()); grouped_dbias.set_rowwise_data(dbias_torch.data_ptr(), GetTransformerEngineDType(tensor.scalar_type()), getTensorShape(dbias_torch)); TensorWrapper workspace_nvte; - auto stream = at::cuda::getCurrentCUDAStream(); + auto stream = getCurrentCUDAStream(); NVTE_SCOPED_GIL_RELEASE({ nvte_group_quantize_dbias(grouped_input_tensor.data(), grouped_output_tensor_cpp.data(), grouped_dbias.data(), workspace_nvte.data(), stream); }); if (workspace_nvte.ndim() > 0 && workspace_nvte.numel() > 0) { - at::Tensor workspace_torch = allocateSpace(workspace_nvte.shape(), workspace_nvte.dtype()); + Tensor workspace_torch = allocateSpace(workspace_nvte.shape(), workspace_nvte.dtype()); workspace_nvte = makeTransformerEngineTensor(workspace_torch.data_ptr(), workspace_nvte.shape(), workspace_nvte.dtype()); } @@ -436,7 +436,7 @@ py::object dequantize(const py::handle &input, transformer_engine::DType otype) auto [out_tensor, out] = q.create_tensor(shape, otype); NVTE_SCOPED_GIL_RELEASE({ - nvte_dequantize(input_tensor.data(), out_tensor.data(), at::cuda::getCurrentCUDAStream()); + nvte_dequantize(input_tensor.data(), out_tensor.data(), getCurrentCUDAStream()); }); return out; @@ -455,10 +455,10 @@ py::object group_dequantize(const py::handle &input, transformer_engine::DType o const auto &quantizer = convert_quantizer(input.attr("quantizer")); // Extract optional tensor attributes. - auto get_optional_tensor = [&input](const char *name) -> std::optional { + auto get_optional_tensor = [&input](const char *name) -> std::optional { auto attr = input.attr(name); if (attr.is_none()) return std::nullopt; - return attr.cast(); + return attr.cast(); }; auto rowwise_data = get_optional_tensor("rowwise_data"); auto columnwise_data = get_optional_tensor("columnwise_data"); @@ -516,7 +516,7 @@ py::object group_dequantize(const py::handle &input, transformer_engine::DType o tensor_offsets, logical_first_dim, logical_last_dim); NVTE_SCOPED_GIL_RELEASE({ - nvte_group_dequantize(input_cpp.data(), out_cpp.data(), at::cuda::getCurrentCUDAStream()); + nvte_group_dequantize(input_cpp.data(), out_cpp.data(), getCurrentCUDAStream()); }); return py::reinterpret_borrow(out_py); @@ -563,7 +563,7 @@ void multi_tensor_quantize_impl(const std::vector &input_list, } NVTE_SCOPED_GIL_RELEASE({ nvte_multi_cast_transpose(nvte_tensor_input_list.size(), nvte_tensor_input_list.data(), - nvte_tensor_output_list.data(), at::cuda::getCurrentCUDAStream()); + nvte_tensor_output_list.data(), getCurrentCUDAStream()); }); } else { // Quantize kernels individually @@ -575,7 +575,7 @@ void multi_tensor_quantize_impl(const std::vector &input_list, } // namespace -std::vector multi_tensor_quantize(const std::vector &tensor_list, +std::vector multi_tensor_quantize(const std::vector &tensor_list, std::vector quantizer_list) { // Check number of tensors const size_t num_tensors = tensor_list.size(); @@ -638,7 +638,7 @@ std::tuple, std::vector> bulk_allocate_fp const auto fp8_dtype = quantizer_cpp_list[0]->dtype; // Allocate row-wise data - std::vector rowwise_data_list, rowwise_scale_list; + std::vector rowwise_data_list, rowwise_scale_list; std::vector> rowwise_data_shapes, rowwise_scale_shapes; if (rowwise_usage) { for (size_t i = 0; i < num_tensors; ++i) { @@ -649,10 +649,10 @@ std::tuple, std::vector> bulk_allocate_fp // Bulk-allocate data and scale tensors std::vector> shapes = rowwise_data_shapes; - std::vector dtypes(num_tensors, torch::kUInt8); + std::vector dtypes(num_tensors, kUInt8); std::vector alignments(num_tensors, 256); shapes.insert(shapes.end(), rowwise_scale_shapes.begin(), rowwise_scale_shapes.end()); - dtypes.insert(dtypes.end(), num_tensors, torch::kFloat32); + dtypes.insert(dtypes.end(), num_tensors, kFloat32); alignments.insert(alignments.end(), num_tensors, 16); auto tensors = bulk_allocate(shapes, dtypes, std::nullopt, alignments); @@ -664,7 +664,7 @@ std::tuple, std::vector> bulk_allocate_fp } // Allocate column-wise data - std::vector columnwise_data_list, columnwise_scale_list; + std::vector columnwise_data_list, columnwise_scale_list; std::vector> columnwise_data_shapes, columnwise_scale_shapes; if (columnwise_usage) { for (size_t i = 0; i < num_tensors; ++i) { @@ -680,10 +680,10 @@ std::tuple, std::vector> bulk_allocate_fp // Bulk-allocate data and scale tensors std::vector> shapes = columnwise_data_shapes; - std::vector dtypes(num_tensors, torch::kUInt8); + std::vector dtypes(num_tensors, kUInt8); std::vector alignments(num_tensors, 256); shapes.insert(shapes.end(), columnwise_scale_shapes.begin(), columnwise_scale_shapes.end()); - dtypes.insert(dtypes.end(), num_tensors, torch::kFloat32); + dtypes.insert(dtypes.end(), num_tensors, kFloat32); alignments.insert(alignments.end(), num_tensors, 16); auto tensors = bulk_allocate(shapes, dtypes, std::nullopt, alignments); @@ -748,7 +748,7 @@ std::tuple, std::vector> bulk_allocate_mx const bool with_gemm_swizzled_scales = quantizer_cpp_list[0]->optimize_for_gemm; // Allocate row-wise data - std::vector rowwise_data_list, rowwise_scale_list; + std::vector rowwise_data_list, rowwise_scale_list; std::vector> rowwise_data_shapes, rowwise_scale_shapes; if (rowwise_usage) { for (size_t i = 0; i < num_tensors; ++i) { @@ -759,10 +759,10 @@ std::tuple, std::vector> bulk_allocate_mx // Bulk-allocate data and scale tensors std::vector> shapes = rowwise_data_shapes; - std::vector dtypes(num_tensors, torch::kUInt8); + std::vector dtypes(num_tensors, kUInt8); std::vector alignments(num_tensors, 256); shapes.insert(shapes.end(), rowwise_scale_shapes.begin(), rowwise_scale_shapes.end()); - dtypes.insert(dtypes.end(), num_tensors, torch::kUInt8); + dtypes.insert(dtypes.end(), num_tensors, kUInt8); alignments.insert(alignments.end(), num_tensors, 16); auto tensors = bulk_allocate(shapes, dtypes, std::nullopt, alignments); @@ -774,7 +774,7 @@ std::tuple, std::vector> bulk_allocate_mx } // Allocate column-wise data - std::vector columnwise_data_list, columnwise_scale_list; + std::vector columnwise_data_list, columnwise_scale_list; std::vector> columnwise_data_shapes, columnwise_scale_shapes; if (columnwise_usage) { for (size_t i = 0; i < num_tensors; ++i) { @@ -787,10 +787,10 @@ std::tuple, std::vector> bulk_allocate_mx // Bulk-allocate data and scale tensors std::vector> shapes = columnwise_data_shapes; - std::vector dtypes(num_tensors, torch::kUInt8); + std::vector dtypes(num_tensors, kUInt8); std::vector alignments(num_tensors, 256); shapes.insert(shapes.end(), columnwise_scale_shapes.begin(), columnwise_scale_shapes.end()); - dtypes.insert(dtypes.end(), num_tensors, torch::kUInt8); + dtypes.insert(dtypes.end(), num_tensors, kUInt8); alignments.insert(alignments.end(), num_tensors, 16); auto tensors = bulk_allocate(shapes, dtypes, std::nullopt, alignments); @@ -919,7 +919,7 @@ std::tuple, std::vector, bool> bulk_alloc }; // Allocate row-wise data - std::vector rowwise_data_list, rowwise_scale_list, amax_rowwise_list; + std::vector rowwise_data_list, rowwise_scale_list, amax_rowwise_list; std::vector> rowwise_data_shapes, rowwise_scale_shapes; if (rowwise_usage) { for (size_t i = 0; i < num_tensors; ++i) { @@ -945,15 +945,15 @@ std::tuple, std::vector, bool> bulk_alloc for (size_t i = 0; i < num_tensors; ++i) { shapes.emplace_back(fp4_byte_shape(rowwise_data_shapes[i])); } - std::vector dtypes(num_tensors, torch::kUInt8); + std::vector dtypes(num_tensors, kUInt8); std::vector alignments(num_tensors, 256); shapes.insert(shapes.end(), rowwise_scale_shapes.begin(), rowwise_scale_shapes.end()); - dtypes.insert(dtypes.end(), num_tensors, torch::kUInt8); + dtypes.insert(dtypes.end(), num_tensors, kUInt8); alignments.insert(alignments.end(), num_tensors, 16); for (size_t i = 0; i < num_tensors; ++i) { shapes.emplace_back(amax_shape(rowwise_data_shapes[i], row_scaled_nvfp4)); } - dtypes.insert(dtypes.end(), num_tensors, torch::kFloat32); + dtypes.insert(dtypes.end(), num_tensors, kFloat32); alignments.insert(alignments.end(), num_tensors, 16); auto tensors = bulk_allocate(shapes, dtypes, std::nullopt, alignments); @@ -966,7 +966,7 @@ std::tuple, std::vector, bool> bulk_alloc } // Allocate column-wise data - std::vector columnwise_data_list, columnwise_scale_list, amax_columnwise_list; + std::vector columnwise_data_list, columnwise_scale_list, amax_columnwise_list; std::vector> columnwise_data_shapes, columnwise_scale_shapes; if (columnwise_usage) { for (size_t i = 0; i < num_tensors; ++i) { @@ -999,15 +999,15 @@ std::tuple, std::vector, bool> bulk_alloc for (size_t i = 0; i < num_tensors; ++i) { shapes.emplace_back(fp4_byte_shape(columnwise_data_shapes[i])); } - std::vector dtypes(num_tensors, torch::kUInt8); + std::vector dtypes(num_tensors, kUInt8); std::vector alignments(num_tensors, 256); shapes.insert(shapes.end(), columnwise_scale_shapes.begin(), columnwise_scale_shapes.end()); - dtypes.insert(dtypes.end(), num_tensors, torch::kUInt8); + dtypes.insert(dtypes.end(), num_tensors, kUInt8); alignments.insert(alignments.end(), num_tensors, 16); for (size_t i = 0; i < num_tensors; ++i) { shapes.emplace_back(amax_shape(columnwise_data_shapes[i])); } - dtypes.insert(dtypes.end(), num_tensors, torch::kFloat32); + dtypes.insert(dtypes.end(), num_tensors, kFloat32); alignments.insert(alignments.end(), num_tensors, 16); auto tensors = bulk_allocate(shapes, dtypes, std::nullopt, alignments); @@ -1077,8 +1077,8 @@ std::tuple, std::vector, bool> bulk_alloc // Owns all allocations/wrappers backing quant_config_list[*].set_rng_state(...). struct StochasticRngStateResources { - at::Tensor rng_states_tensor; // [2 * num_tensors], int64, CUDA - at::Tensor rng_states_tensor_colwise; // optional, same shape/dtype/device + Tensor rng_states_tensor; // [2 * num_tensors], int64, CUDA + Tensor rng_states_tensor_colwise; // optional, same shape/dtype/device std::vector te_rng_state_list; std::vector te_rng_state_list_colwise; @@ -1113,21 +1113,21 @@ static StochasticRngStateResources setup_stochastic_rounding_rng_states_helper( const size_t rng_elts_per_thread = res.with_bulk_generate_rng_states ? (1024 * num_tensors) : 1024; - auto opts = at::TensorOptions().dtype(torch::kInt64).device(torch::kCUDA); - res.rng_states_tensor = torch::empty({static_cast(2 * num_tensors)}, opts); + auto opts = TensorOptions().dtype(kInt64).device(kCUDA); + res.rng_states_tensor = empty({static_cast(2 * num_tensors)}, opts); if (need_separate_rng_states) { - res.rng_states_tensor_colwise = torch::empty({static_cast(2 * num_tensors)}, opts); + res.rng_states_tensor_colwise = empty({static_cast(2 * num_tensors)}, opts); } res.te_rng_state_list.reserve(num_tensors); if (need_separate_rng_states) res.te_rng_state_list_colwise.reserve(num_tensors); for (size_t i = 0; i < num_tensors; ++i) { - auto gen = at::get_generator_or_default( - std::nullopt, at::cuda::detail::getDefaultCUDAGenerator()); + auto gen = get_generator_or_default( + std::nullopt, getDefaultCUDAGenerator()); // Rowwise RNG state - at::PhiloxCudaState philox_args = init_philox_state(gen, rng_elts_per_thread); + PhiloxCudaState philox_args = init_philox_state(gen, rng_elts_per_thread); int64_t *rng_state_ptr = static_cast(res.rng_states_tensor.data_ptr()) + i * 2; philox_unpack(philox_args, rng_state_ptr); @@ -1139,7 +1139,7 @@ static StochasticRngStateResources setup_stochastic_rounding_rng_states_helper( // Colwise RNG state (only if you truly need a different sequence) if (need_separate_rng_states) { // re-initialize philox_args for colwise RNG state - at::PhiloxCudaState philox_args_col = init_philox_state(gen, rng_elts_per_thread); + PhiloxCudaState philox_args_col = init_philox_state(gen, rng_elts_per_thread); int64_t *rng_state_ptr_colwise = static_cast(res.rng_states_tensor_colwise.data_ptr()) + i * 2; @@ -1266,7 +1266,7 @@ void split_quantize_nvfp4_impl_with_rht_helper(const TensorWrapper &input, if (all_aligned_token_dim) { // allocate a tile scheduler workspace auto tile_scheduler_workspace_torch = - at::empty({1}, at::device(at::kCUDA).dtype(torch::kInt32)); + empty({1}, TensorOptions().device(kCUDA).dtype(kInt32)); auto nvte_tile_scheduler_workspace = makeTransformerEngineTensor(tile_scheduler_workspace_torch); // call the fully-fused grouped kernel for rowwise quantization & colwise RHT quantization transpose @@ -1488,7 +1488,7 @@ void split_quantize_nvfp4_impl(const TensorWrapper &input, "NVFP4 multi-quantize requires inner dim to be multiple of 128."); // CUDA stream - auto stream = at::cuda::getCurrentCUDAStream(); + auto stream = getCurrentCUDAStream(); // Perform multi-tensor quantization NVTE_SCOPED_GIL_RELEASE({ @@ -1508,7 +1508,7 @@ void split_quantize_nvfp4_impl(const TensorWrapper &input, } // namespace -std::vector split_quantize(const at::Tensor &tensor, +std::vector split_quantize(const Tensor &tensor, const std::vector &split_sections, std::vector quantizer_list, bool disable_bulk_allocation) { diff --git a/transformer_engine/pytorch/csrc/extensions/comm_gemm_overlap.cpp b/transformer_engine/pytorch/csrc/extensions/comm_gemm_overlap.cpp index 33237f0751..92cd6321b1 100644 --- a/transformer_engine/pytorch/csrc/extensions/comm_gemm_overlap.cpp +++ b/transformer_engine/pytorch/csrc/extensions/comm_gemm_overlap.cpp @@ -13,7 +13,10 @@ #define HALF_BYTES 2 #define UB_MAX_SM 32 -using namespace torch::indexing; +// CommOverlap methods and helpers in this TU live in the global namespace; +// bring the facade aliases/functions into scope (also makes `indexing` visible). +using namespace transformer_engine::pytorch; // NOLINT(build/namespaces) +using namespace indexing; using namespace std::placeholders; namespace te = transformer_engine; @@ -28,14 +31,14 @@ CommOverlapHelper::CommOverlapHelper() { #endif } // empty constructor for NVTE_UB_WITH_MPI=1 -CommOverlapHelper::CommOverlapHelper(c10d::ProcessGroup *world_group, - std::optional intra_domain_group) { +CommOverlapHelper::CommOverlapHelper(ProcessGroup *world_group, + std::optional intra_domain_group) { #ifndef NVTE_UB_WITH_MPI torch_pgs.insert({"world", world_group}); myrank = torch_pgs["world"]->getRank(); numranks = torch_pgs["world"]->getSize(); - c10d::ProcessGroup::BackendType backend = torch_pgs["world"]->getBackendType(); - backend_is_nccl = (backend == c10d::ProcessGroup::BackendType::NCCL); + ProcessGroup::BackendType backend = torch_pgs["world"]->getBackendType(); + backend_is_nccl = (backend == ProcessGroup::BackendType::NCCL); if (intra_domain_group.has_value()) { // Get local rank on node and number of local ranks @@ -78,13 +81,13 @@ CommOverlapHelper::CommOverlapHelper(c10d::ProcessGroup *world_group, NVTE_CHECK_NCCL(ncclGetUniqueId(&nccl_world_id)); } auto nccl_world_id_tensor = - torch::from_blob(reinterpret_cast(&nccl_world_id), {sizeof(ncclUniqueId)}, - at::device(torch::kCPU).dtype(torch::kUInt8)); + from_blob(reinterpret_cast(&nccl_world_id), {sizeof(ncclUniqueId)}, + TensorOptions().device(kCPU).dtype(kUInt8)); nccl_world_id_tensor = (backend_is_nccl) ? nccl_world_id_tensor.cuda() : nccl_world_id_tensor; { - c10d::BroadcastOptions bcast_opts; + BroadcastOptions bcast_opts; bcast_opts.rootRank = 0; - std::vector bcast_tensors = {nccl_world_id_tensor}; + std::vector bcast_tensors = {nccl_world_id_tensor}; auto work = torch_pgs["world"]->broadcast(bcast_tensors, bcast_opts); work->wait(); } @@ -104,13 +107,13 @@ CommOverlapHelper::CommOverlapHelper(c10d::ProcessGroup *world_group, // Broadcast the intra-node unique ID from the local root to all local ranks auto nccl_intra_id_tensor = - torch::from_blob(reinterpret_cast(&nccl_intra_id), {sizeof(ncclUniqueId)}, - at::device(torch::kCPU).dtype(torch::kUInt8)); + from_blob(reinterpret_cast(&nccl_intra_id), {sizeof(ncclUniqueId)}, + TensorOptions().device(kCPU).dtype(kUInt8)); nccl_intra_id_tensor = (backend_is_nccl) ? nccl_intra_id_tensor.cuda() : nccl_intra_id_tensor; { - c10d::BroadcastOptions bcast_opts; + BroadcastOptions bcast_opts; bcast_opts.rootRank = 0; - std::vector bcast_tensors = {nccl_intra_id_tensor}; + std::vector bcast_tensors = {nccl_intra_id_tensor}; auto work = torch_pgs["intra"]->broadcast(bcast_tensors, bcast_opts); work->wait(); } @@ -155,24 +158,24 @@ void CommOverlapHelper::ub_allgather(void *globaldata, size_t globalbytes, void "with valid process groups!"); auto localtensor = - torch::from_blob(localdata, {static_cast(localbytes / sizeof(uint8_t))}, - at::device(torch::kCPU).dtype(torch::kUInt8)); + from_blob(localdata, {static_cast(localbytes / sizeof(uint8_t))}, + TensorOptions().device(kCPU).dtype(kUInt8)); auto localtmp = (backend_is_nccl) ? localtensor.cuda() : localtensor; auto globaltensor = - torch::from_blob(globaldata, {static_cast(globalbytes / sizeof(uint8_t))}, - at::device(torch::kCPU).dtype(torch::kUInt8)); + from_blob(globaldata, {static_cast(globalbytes / sizeof(uint8_t))}, + TensorOptions().device(kCPU).dtype(kUInt8)); auto globaltmp = (backend_is_nccl) ? globaltensor.cuda() : globaltensor; - std::vector> globalchunks = { + std::vector> globalchunks = { globaltmp.chunk(torch_pgs[group]->getSize())}; - std::vector localchunk = {localtmp}; + std::vector localchunk = {localtmp}; auto work = torch_pgs[group]->allgather(globalchunks, localchunk); work->wait(); if (backend_is_nccl) { globaltensor.copy_(globaltmp.cpu()); - globaltmp = torch::Tensor(); - localtmp = torch::Tensor(); + globaltmp = Tensor(); + localtmp = Tensor(); } #else NVTE_ERROR("Internal TE error: CommOverlapHelper::ub_allgather is a no-op when TE is compiled ", @@ -213,7 +216,7 @@ CommOverlapHelper::NcclCommSharedPtr CommOverlapHelper::get_nccl_comm(std::strin * CommOverlap **************************************************************************************************/ -CommOverlap::CommOverlap(const std::vector &buffer_shape, at::ScalarType buffer_dtype, +CommOverlap::CommOverlap(const std::vector &buffer_shape, ScalarType buffer_dtype, CommOverlapHelper *helper, int tp_size, int num_splits, int num_max_streams, int comm_cga_size, int gemm_priority, int comm_priority, int num_comm_sm, bool set_sm_margin, bool atomic_gemm, @@ -279,7 +282,7 @@ void cublasmp_capture_warmup(te::CommOverlapCore *core, int tp_size, te::CommOve B_tw.set_rowwise_data(b_ptr, te::DType::kBFloat16, b_shape); D_tw.set_rowwise_data(d_ptr, te::DType::kBFloat16, d_shape); - cudaStream_t stream = at::cuda::getCurrentCUDAStream(); + cudaStream_t stream = getCurrentCUDAStream(); if (comm_type == te::CommOverlapType::AG) { if (core->is_atomic_gemm()) { core->atomic_gemm_overlap_ag( @@ -311,7 +314,7 @@ void cublasmp_capture_warmup(te::CommOverlapCore *core, int tp_size, te::CommOve CommOverlap::CommOverlap(CommOverlapHelper *helper, int tp_rank, int tp_size, te::CommOverlapType comm_type, const std::vector &buffer_shape, - at::ScalarType buffer_dtype, int num_comm_sm, bool atomic_gemm) + ScalarType buffer_dtype, int num_comm_sm, bool atomic_gemm) : te::CommOverlapBase(helper->get_nccl_comm("intra").get(), tp_rank, tp_size, num_comm_sm, atomic_gemm), _nccl_comm(helper->get_nccl_comm("intra")) { @@ -324,7 +327,7 @@ CommOverlap::CommOverlap(CommOverlapHelper *helper, int tp_rank, int tp_size, /* ** Helper function to copy input to _ubuf */ -void CommOverlap::copy_into_buffer(const at::Tensor &input, bool local_chunk) { +void CommOverlap::copy_into_buffer(const Tensor &input, bool local_chunk) { const auto &input_ = input.contiguous(); // Check element size @@ -354,14 +357,14 @@ void CommOverlap::copy_into_buffer(const at::Tensor &input, bool local_chunk) { } // Copy data - auto stream_main = at::cuda::getCurrentCUDAStream(); + auto stream_main = getCurrentCUDAStream(); NVTE_CHECK_CUDA(cudaEventRecord(_start_d2dcopy, (cudaStream_t)stream_main)); NVTE_CHECK_CUDA(cudaStreamWaitEvent((cudaStream_t)_stream_comm, _start_d2dcopy, 0)); NVTE_CHECK_CUDA(cudaMemcpyAsync(dst_ptr, src_ptr, input_size * element_size, cudaMemcpyDeviceToDevice, (cudaStream_t)_stream_comm)); } -at::Tensor CommOverlap::get_buffer(bool local_chunk, std::optional> shape) { +Tensor CommOverlap::get_buffer(bool local_chunk, std::optional> shape) { // Check buffer shape const size_t ubuf_size = _ubuf.numel(); if (shape) { @@ -393,20 +396,20 @@ at::Tensor CommOverlap::get_buffer(bool local_chunk, std::optional CommOverlap::get_communication_stream() { +std::pair CommOverlap::get_communication_stream() { // Return the same stream for both send and recv - return {at::cuda::getStreamFromExternal(_stream_comm, at::cuda::current_device()), - at::cuda::getStreamFromExternal(_stream_comm, at::cuda::current_device())}; + return {getStreamFromExternal(_stream_comm, current_device()), + getStreamFromExternal(_stream_comm, current_device())}; } /*************************************************************************************************** * CommOverlapP2P **************************************************************************************************/ -CommOverlapP2P::CommOverlapP2P(const std::vector &buffer_shape, at::ScalarType buffer_dtype, +CommOverlapP2P::CommOverlapP2P(const std::vector &buffer_shape, ScalarType buffer_dtype, CommOverlapHelper *helper, int tp_size, te::CommOverlapType comm_type, int num_max_streams, int comm_cga_size, int gemm_priority, int comm_priority, @@ -422,7 +425,7 @@ CommOverlapP2P::CommOverlapP2P(const std::vector &buffer_shape, at::Scal CommOverlapP2P::CommOverlapP2P(CommOverlapHelper *helper, int tp_rank, int tp_size, te::CommOverlapType comm_type, - const std::vector &buffer_shape, at::ScalarType buffer_dtype, + const std::vector &buffer_shape, ScalarType buffer_dtype, int num_comm_sm, bool atomic_gemm) : te::CommOverlapP2PBase(helper->get_nccl_comm("intra").get(), tp_rank, tp_size, num_comm_sm, atomic_gemm), @@ -435,7 +438,7 @@ CommOverlapP2P::CommOverlapP2P(CommOverlapHelper *helper, int tp_rank, int tp_si /* ** Copy input to _ubufs[0] */ -void CommOverlapP2P::copy_into_buffer(const at::Tensor &input, bool local_chunk) { +void CommOverlapP2P::copy_into_buffer(const Tensor &input, bool local_chunk) { const auto &input_ = input.contiguous(); // Check element size @@ -466,10 +469,10 @@ void CommOverlapP2P::copy_into_buffer(const at::Tensor &input, bool local_chunk) // Copy data NVTE_CHECK_CUDA(cudaMemcpyAsync(dst_ptr, src_ptr, input_size * element_size, cudaMemcpyDeviceToDevice, - (cudaStream_t)at::cuda::getCurrentCUDAStream())); + (cudaStream_t)getCurrentCUDAStream())); } -at::Tensor CommOverlapP2P::get_buffer(bool local_chunk, std::optional> shape) { +Tensor CommOverlapP2P::get_buffer(bool local_chunk, std::optional> shape) { // Check buffer shape if (shape) { const size_t requested_size = transformer_engine::pytorch::product(*shape); @@ -496,17 +499,17 @@ at::Tensor CommOverlapP2P::get_buffer(bool local_chunk, std::optional CommOverlapP2P::get_communication_stream() { - return {at::cuda::getStreamFromExternal(_stream_send[0], at::cuda::current_device()), - at::cuda::getStreamFromExternal(_stream_recv, at::cuda::current_device())}; +std::pair CommOverlapP2P::get_communication_stream() { + return {getStreamFromExternal(_stream_send[0], current_device()), + getStreamFromExternal(_stream_recv, current_device())}; } void transformer_engine::pytorch::bulk_overlap_ag_with_external_gemm( - CommOverlap &allgather_communicator, at::Stream send_stream, at::Stream recv_stream) { - auto main_stream = at::cuda::getCurrentCUDAStream(); - allgather_communicator.bulk_overlap_external_ag(at::cuda::CUDAStream(send_stream), - at::cuda::CUDAStream(recv_stream), main_stream); + CommOverlap &allgather_communicator, Stream send_stream, Stream recv_stream) { + auto main_stream = getCurrentCUDAStream(); + allgather_communicator.bulk_overlap_external_ag(CUDAStream(send_stream), + CUDAStream(recv_stream), main_stream); } diff --git a/transformer_engine/pytorch/csrc/extensions/dropout.cpp b/transformer_engine/pytorch/csrc/extensions/dropout.cpp index bea8f3a7b5..c46d22bbec 100644 --- a/transformer_engine/pytorch/csrc/extensions/dropout.cpp +++ b/transformer_engine/pytorch/csrc/extensions/dropout.cpp @@ -6,11 +6,8 @@ #include "transformer_engine/dropout.h" -#include #include -#include - #include "../common.h" #include "../extensions.h" #include "../pybind.h" @@ -20,7 +17,7 @@ namespace transformer_engine { namespace pytorch { std::vector dropout_fwd(const py::handle &input, float dropout_probability, - std::optional out) { + std::optional out) { using namespace transformer_engine::pytorch::detail; // Input tensor @@ -28,14 +25,14 @@ std::vector dropout_fwd(const py::handle &input, float dropout_proba // Allocate output tensor if needed if (!out) { - at::ScalarType dtype = GetATenDType(input_nvte.dtype()); - if (dtype == at::kFloat8_e4m3fn || dtype == at::kFloat8_e5m2) { - dtype = input.attr("dtype").cast(); + ScalarType dtype = GetATenDType(input_nvte.dtype()); + if (dtype == kFloat8_e4m3fn || dtype == kFloat8_e5m2) { + dtype = input.attr("dtype").cast(); } const auto shape_uint64 = convertShape(input_nvte.shape()); const std::vector shape_int64(shape_uint64.begin(), shape_uint64.end()); - const auto opts = at::TensorOptions().dtype(dtype).device(torch::kCUDA); - out = at::empty(shape_int64, opts); + const auto opts = TensorOptions().dtype(dtype).device(kCUDA); + out = empty(shape_int64, opts); } TensorWrapper out_nvte = makeTransformerEngineTensor(*out); @@ -44,9 +41,9 @@ std::vector dropout_fwd(const py::handle &input, float dropout_proba auto mask_nvte = makeTransformerEngineTensor(mask_pyt); // RNG state tensor - auto gen = at::get_generator_or_default( - std::nullopt, at::cuda::detail::getDefaultCUDAGenerator()); - at::PhiloxCudaState philox_args; + auto gen = get_generator_or_default( + std::nullopt, getDefaultCUDAGenerator()); + PhiloxCudaState philox_args; { std::lock_guard lock(gen->mutex_); constexpr int64_t rng_elts_per_thread = 4; @@ -57,30 +54,30 @@ std::vector dropout_fwd(const py::handle &input, float dropout_proba nvte_extract_seed_and_offset( reinterpret_cast(rng_state_pyt.data_ptr()), philox_args.captured_, philox_args.seed_.ptr, philox_args.seed_.val, philox_args.offset_.ptr, - philox_args.offset_.val, philox_args.offset_intragraph_, at::cuda::getCurrentCUDAStream()); + philox_args.offset_.val, philox_args.offset_intragraph_, getCurrentCUDAStream()); }); auto rng_state_nvte = makeTransformerEngineTensor(rng_state_pyt); // Launch kernel NVTE_SCOPED_GIL_RELEASE({ nvte_dropout_fwd(input_nvte.data(), out_nvte.data(), mask_nvte.data(), rng_state_nvte.data(), - dropout_probability, at::cuda::getCurrentCUDAStream()); + dropout_probability, getCurrentCUDAStream()); }); return {py::cast(std::move(*out)), py::cast(mask_pyt)}; } -py::object dropout_bwd(const at::Tensor &grad_output, const at::Tensor &mask, - const float dropout_probability, std::optional grad_input) { +py::object dropout_bwd(const Tensor &grad_output, const Tensor &mask, + const float dropout_probability, std::optional grad_input) { const auto grad_output_nvte = makeTransformerEngineTensor(grad_output); const auto mask_nvte = makeTransformerEngineTensor(mask); if (!grad_input) { - grad_input = at::empty_like(grad_output); + grad_input = empty_like(grad_output); } auto grad_input_nvte = makeTransformerEngineTensor(*grad_input); NVTE_SCOPED_GIL_RELEASE({ nvte_dropout_bwd(grad_output_nvte.data(), mask_nvte.data(), grad_input_nvte.data(), - dropout_probability, at::cuda::getCurrentCUDAStream()); + dropout_probability, getCurrentCUDAStream()); }); return py::cast(std::move(*grad_input)); } diff --git a/transformer_engine/pytorch/csrc/extensions/fp8_partial_cast.cpp b/transformer_engine/pytorch/csrc/extensions/fp8_partial_cast.cpp index d6693a485e..b43f448a4d 100644 --- a/transformer_engine/pytorch/csrc/extensions/fp8_partial_cast.cpp +++ b/transformer_engine/pytorch/csrc/extensions/fp8_partial_cast.cpp @@ -8,13 +8,13 @@ namespace transformer_engine::pytorch { -void fp8_block_scaling_compute_partial_amax(const at::Tensor &tensor, at::Tensor amax, size_t h, +void fp8_block_scaling_compute_partial_amax(const Tensor &tensor, Tensor amax, size_t h, size_t w, size_t start_offset, size_t block_len) { TORCH_CHECK(block_len == 128, "Currently only block_len = 128 is supported"); TORCH_CHECK(amax.dim() == 2, "amax must be a 2D tensor"); - TORCH_CHECK(amax.scalar_type() == at::ScalarType::Float, "amax must be a float tensor"); - TORCH_CHECK(tensor.scalar_type() == at::ScalarType::Float || - tensor.scalar_type() == at::ScalarType::BFloat16, + TORCH_CHECK(amax.scalar_type() == ScalarType::Float, "amax must be a float tensor"); + TORCH_CHECK(tensor.scalar_type() == ScalarType::Float || + tensor.scalar_type() == ScalarType::BFloat16, "tensor must be a float or bfloat16 tensor"); const TensorWrapper tensor_cu = makeTransformerEngineTensor(tensor); @@ -22,19 +22,19 @@ void fp8_block_scaling_compute_partial_amax(const at::Tensor &tensor, at::Tensor nvte_fp8_block_scaling_compute_partial_amax(tensor_cu.data(), amax_cu.data(), h, w, amax.stride(0), amax.stride(1), start_offset, - block_len, at::cuda::getCurrentCUDAStream()); + block_len, getCurrentCUDAStream()); } -void fp8_block_scaling_partial_cast(const at::Tensor &inp, at::Tensor out, const at::Tensor &scale, +void fp8_block_scaling_partial_cast(const Tensor &inp, Tensor out, const Tensor &scale, size_t h, size_t w, size_t start_offset, size_t block_len, const transformer_engine::DType out_dtype) { TORCH_CHECK(block_len == 128, "Currently only block_len = 128 is supported"); TORCH_CHECK(scale.dim() == 2, "scale must be a 2D tensor"); - TORCH_CHECK(scale.scalar_type() == at::ScalarType::Float, "scale must be a float tensor"); + TORCH_CHECK(scale.scalar_type() == ScalarType::Float, "scale must be a float tensor"); TORCH_CHECK( - inp.scalar_type() == at::ScalarType::Float || inp.scalar_type() == at::ScalarType::BFloat16, + inp.scalar_type() == ScalarType::Float || inp.scalar_type() == ScalarType::BFloat16, "input must be a float or bfloat16 tensor"); - TORCH_CHECK(out.scalar_type() == at::ScalarType::Byte, "output must be a uint8 tensor"); + TORCH_CHECK(out.scalar_type() == ScalarType::Byte, "output must be a uint8 tensor"); TORCH_CHECK(out_dtype == transformer_engine::DType::kFloat8E4M3 || out_dtype == transformer_engine::DType::kFloat8E5M2, "out_dtype must be kFloat8E4M3 or kFloat8E5M2"); @@ -45,11 +45,11 @@ void fp8_block_scaling_partial_cast(const at::Tensor &inp, at::Tensor out, const nvte_fp8_block_scaling_partial_cast( inp_cu.data(), out_cu.data(), scale_cu.data(), h, w, scale.stride(0), scale.stride(1), - start_offset, block_len, static_cast(out_dtype), at::cuda::getCurrentCUDAStream()); + start_offset, block_len, static_cast(out_dtype), getCurrentCUDAStream()); } -void mxfp8_scaling_compute_partial_amax(const at::Tensor &input, at::Tensor amax_rowwise, - at::Tensor amax_colwise, int rows, int cols, +void mxfp8_scaling_compute_partial_amax(const Tensor &input, Tensor amax_rowwise, + Tensor amax_colwise, int rows, int cols, size_t start_offset) { TORCH_CHECK(input.is_contiguous(), "input must be contiguous"); TORCH_CHECK(amax_rowwise.is_contiguous(), "amax_rowwise must be contiguous"); @@ -61,12 +61,12 @@ void mxfp8_scaling_compute_partial_amax(const at::Tensor &input, at::Tensor amax nvte_mxfp8_scaling_compute_partial_amax(input_cu.data(), amax_rowwise_cu.data(), amax_colwise_cu.data(), rows, cols, start_offset, - at::cuda::getCurrentCUDAStream()); + getCurrentCUDAStream()); } -void mxfp8_scaling_partial_cast(const at::Tensor &input, at::Tensor output_rowwise, - at::Tensor output_colwise, const at::Tensor &scale_inv_rowwise, - const at::Tensor &scale_inv_colwise, int rows, int cols, +void mxfp8_scaling_partial_cast(const Tensor &input, Tensor output_rowwise, + Tensor output_colwise, const Tensor &scale_inv_rowwise, + const Tensor &scale_inv_colwise, int rows, int cols, size_t start_offset) { TORCH_CHECK(input.is_contiguous(), "input must be contiguous"); TORCH_CHECK(output_rowwise.is_contiguous(), "output_rowwise must be contiguous"); @@ -83,7 +83,7 @@ void mxfp8_scaling_partial_cast(const at::Tensor &input, at::Tensor output_rowwi nvte_mxfp8_scaling_partial_cast(input_cu.data(), output_rowwise_cu.data(), output_colwise_cu.data(), scale_inv_rowwise_cu.data(), scale_inv_colwise_cu.data(), rows, cols, start_offset, - at::cuda::getCurrentCUDAStream()); + getCurrentCUDAStream()); } } // namespace transformer_engine::pytorch diff --git a/transformer_engine/pytorch/csrc/extensions/gemm.cpp b/transformer_engine/pytorch/csrc/extensions/gemm.cpp index b1e552ec8b..c9728d5ff9 100644 --- a/transformer_engine/pytorch/csrc/extensions/gemm.cpp +++ b/transformer_engine/pytorch/csrc/extensions/gemm.cpp @@ -98,9 +98,9 @@ struct GroupedGemmConfig { std::optional matmul_config; }; -GroupedGemmConfig prepare_grouped_gemm_config(at::Tensor alpha, at::Tensor beta, - at::Tensor workspace_setup, - at::Tensor workspace_cublas, size_t num_tensors, +GroupedGemmConfig prepare_grouped_gemm_config(Tensor alpha, Tensor beta, + Tensor workspace_setup, + Tensor workspace_cublas, size_t num_tensors, int math_sm_count, bool use_split_accumulator) { const bool per_group = (alpha.numel() == static_cast(num_tensors)); const bool scalar = (alpha.numel() == 1); @@ -144,7 +144,7 @@ std::pair createOutputTensor(const std::vector gemm(py::handle A, bool transa, py::handle B, bool transb, py::object D, py::handle quantizer, std::optional out_dtype, MaybeTensor bias, DType bias_type, bool gelu, MaybeTensor gelu_in, bool grad, - at::Tensor workspace, size_t workspaceSize, bool accumulate, + Tensor workspace, size_t workspaceSize, bool accumulate, bool use_split_accumulator, CommOverlapCore* comm_overlap, std::optional comm_type, MaybeTensor extra_output, bool bulk_overlap, float alpha, std::optional beta) { @@ -153,7 +153,7 @@ std::vector gemm(py::handle A, bool transa, py::handle B, bool trans // Ensure that cublasLt handle is created on the correct device, // overriding torch.cuda.set_device calls from user side. // Assumes all tensors passed are on the same device. - at::cuda::CUDAGuard device_guard(workspace.device()); + CUDAGuard device_guard(workspace.device()); // Input tensors NVTE_CHECK(!A.is_none(), "Tensor A has not been provided"); @@ -240,8 +240,8 @@ std::vector gemm(py::handle A, bool transa, py::handle B, bool trans if (bias.has_value()) { if (grad) { auto opts = - torch::TensorOptions().dtype(GetATenDType(out_tensor.dtype())).device(torch::kCUDA); - bias_grad = at::empty({static_cast(B_shape.data[B_shape.ndim - 1])}, opts); + TensorOptions().dtype(GetATenDType(out_tensor.dtype())).device(kCUDA); + bias_grad = empty({static_cast(B_shape.data[B_shape.ndim - 1])}, opts); bias_tensor = makeTransformerEngineTensor(*bias_grad); } else { if (!bias->is_contiguous()) { @@ -257,12 +257,12 @@ std::vector gemm(py::handle A, bool transa, py::handle B, bool trans if (gelu) { if (!grad) { auto dtype = GetATenDType(gelu_type); - auto opts = torch::TensorOptions().dtype(dtype).device(torch::kCUDA); + auto opts = TensorOptions().dtype(dtype).device(kCUDA); std::vector torch_shape; for (auto v : D_shape) { torch_shape.push_back(v); } - pre_gelu_out = at::empty(torch_shape, opts); + pre_gelu_out = empty(torch_shape, opts); } else { if (gelu_in.has_value()) { pre_gelu_out = *gelu_in; @@ -280,7 +280,7 @@ std::vector gemm(py::handle A, bool transa, py::handle B, bool trans // Set an external SM Margin to all the GEMMs. // This comes in handy when DP is overlapped with GEMMs - const int device_id = at::cuda::current_device(); + const int device_id = current_device(); const int sm_count = transformer_engine::cuda::sm_count(device_id); int num_math_sms = sm_count - transformer_engine::getenv("NVTE_EXT_MARGIN_SM", sm_count); @@ -298,8 +298,8 @@ std::vector gemm(py::handle A, bool transa, py::handle B, bool trans config.set_sm_count(num_math_sms); // Keep the swizzled scaling factor tensors alive during the GEMM. - std::vector> swizzled_scale_inverses_list; - auto main_stream = at::cuda::getCurrentCUDAStream(); + std::vector> swizzled_scale_inverses_list; + auto main_stream = getCurrentCUDAStream(); if (A_tensor.numel() != 0 && B_tensor.numel() != 0) { // Optionally swizzle the scaling factors auto [A_row_scales, A_col_scales] = swizzle_scales_for_gemm(A_tensor, transa, !transa); @@ -413,18 +413,18 @@ std::vector gemm(py::handle A, bool transa, py::handle B, bool trans return out; } -void te_atomic_gemm(at::Tensor A, at::Tensor A_scale_inverse, DType A_type, - std::vector A_scaling_mode, bool transa, at::Tensor B, - at::Tensor B_scale_inverse, DType B_type, std::vector B_scaling_mode, - bool transb, at::Tensor D, at::Tensor D_scale, DType D_type, at::Tensor D_amax, - at::Tensor bias, DType bias_type, at::Tensor pre_gelu_out, bool grad, - at::Tensor workspace, size_t workspaceSize, bool accumulate, +void te_atomic_gemm(Tensor A, Tensor A_scale_inverse, DType A_type, + std::vector A_scaling_mode, bool transa, Tensor B, + Tensor B_scale_inverse, DType B_type, std::vector B_scaling_mode, + bool transb, Tensor D, Tensor D_scale, DType D_type, Tensor D_amax, + Tensor bias, DType bias_type, Tensor pre_gelu_out, bool grad, + Tensor workspace, size_t workspaceSize, bool accumulate, bool use_split_accumulator, int math_sm_count, int m_split, int n_split, - bool gemm_producer, at::Tensor counter) { + bool gemm_producer, Tensor counter) { // Ensure that cublasLt handle is created on the correct device, // overriding torch.cuda.set_device calls from user side. // Assumes all tensors passed are on the same device. - at::cuda::CUDAGuard device_guard(workspace.device()); + CUDAGuard device_guard(workspace.device()); // TODO: Handle scaling modes NVTEScalingMode nvte_scaling_modeA = NVTE_DELAYED_TENSOR_SCALING; @@ -461,15 +461,15 @@ void te_atomic_gemm(at::Tensor A, at::Tensor A_scale_inverse, DType A_type, nvte_cublas_atomic_gemm(te_A.data(), te_B.data(), te_D.data(), te_bias.data(), te_pre_gelu_out.data(), transa, transb, grad, te_workspace.data(), accumulate, use_split_accumulator, math_sm_count, m_split, n_split, - gemm_producer, te_counter.data(), at::cuda::getCurrentCUDAStream()); + gemm_producer, te_counter.data(), getCurrentCUDAStream()); }); } -std::optional> te_general_grouped_gemm( +std::optional> te_general_grouped_gemm( std::vector A, bool transa, std::vector B, bool transb, - std::optional> D, DType D_type, std::vector m_splits, - std::vector bias, DType bias_type, bool single_output, - std::vector pre_gelu_out, bool grad, std::vector workspace, + std::optional> D, DType D_type, std::vector m_splits, + std::vector bias, DType bias_type, bool single_output, + std::vector pre_gelu_out, bool grad, std::vector workspace, size_t workspaceSize, bool accumulate, bool use_split_accumulator, int math_sm_count) { if (single_output && D == std::nullopt) { NVTE_ERROR("not implemented, D should be allocated for single output case."); @@ -478,7 +478,7 @@ std::optional> te_general_grouped_gemm( // Ensure that cublasLt handle is created on the correct device, // overriding torch.cuda.set_device calls from user side. // Assumes all tensors passed are on the same device. - at::cuda::CUDAGuard device_guard(workspace[0].device()); + CUDAGuard device_guard(workspace[0].device()); void* output_data_ptr = nullptr; if (single_output) { @@ -488,13 +488,13 @@ std::optional> te_general_grouped_gemm( const auto none = py::none(); std::vector te_A_wrappers, te_B_wrappers, te_D_wrappers, te_bias_wrappers, te_pre_gelu_out_wrappers; - std::vector D_vectors; + std::vector D_vectors; for (size_t i = 0; i < A.size(); i++) { auto te_A = makeTransformerEngineTensor(A[i], none); auto te_B = makeTransformerEngineTensor(B[i], none); // if there is single output - at::Tensor out_tensor; + Tensor out_tensor; auto size_t_shape = pytorch::detail::getGemmOutputShape(te_A.shape(), transa, te_B.shape(), transb); bool D_numel_is_zero = false; @@ -506,16 +506,16 @@ std::optional> te_general_grouped_gemm( } } auto dtype = GetATenDType(D_type); - auto opts = torch::TensorOptions().dtype(dtype).device(torch::kCUDA); + auto opts = TensorOptions().dtype(dtype).device(kCUDA); if (single_output) { if (output_data_ptr == nullptr) { - out_tensor = at::empty(D_shape, opts); + out_tensor = empty(D_shape, opts); } else { // We need to check !D_numel_is_zero because if the final input portion has zero elements, // output_data_ptr would point beyond the allocated memory of D. This would cause - // at::from_blob to fail as it would reference memory not allocated by CUDA. + // from_blob to fail as it would reference memory not allocated by CUDA. if (!D_numel_is_zero) { - out_tensor = at::from_blob(output_data_ptr, D_shape, opts); + out_tensor = from_blob(output_data_ptr, D_shape, opts); } } char* char_ptr = reinterpret_cast(output_data_ptr); @@ -524,8 +524,8 @@ std::optional> te_general_grouped_gemm( D_vectors.emplace_back(out_tensor); } else { if (D == std::nullopt) { - auto opts = torch::TensorOptions().dtype(dtype).device(torch::kCUDA); - out_tensor = at::empty(D_shape, opts); + auto opts = TensorOptions().dtype(dtype).device(kCUDA); + out_tensor = empty(D_shape, opts); D_vectors.emplace_back(out_tensor); } else { out_tensor = (*D)[i]; @@ -566,7 +566,7 @@ std::optional> te_general_grouped_gemm( } // Keep the swizzled scaling factor tensors alive during the GEMM. - std::vector> swizzled_scale_inverses_list; + std::vector> swizzled_scale_inverses_list; // Optionally swizzle the scaling factors swizzled_scale_inverses_list.emplace_back( @@ -631,15 +631,15 @@ std::optional> te_general_grouped_gemm( nvte_multi_tensor_gemm(te_A_vector.data(), te_B_vector.data(), te_D_vector.data(), te_bias_vector.data(), te_pre_gelu_out_vector.data(), te_A_vector.size(), transa, transb, grad, te_workspace_vector.data(), accumulate, - use_split_accumulator, math_sm_count, at::cuda::getCurrentCUDAStream()); + use_split_accumulator, math_sm_count, getCurrentCUDAStream()); }); return bias; } py::object te_general_grouped_gemm_for_grouped_tensor( py::handle A, bool transa, py::handle B, bool transb, py::handle D, py::object bias, - std::optional bias_scale, at::Tensor alpha, at::Tensor beta, - at::Tensor workspace_setup, at::Tensor workspace_cublas, bool use_split_accumulator, + std::optional bias_scale, Tensor alpha, Tensor beta, + Tensor workspace_setup, Tensor workspace_cublas, bool use_split_accumulator, int math_sm_count) { using namespace transformer_engine::pytorch::detail; @@ -648,7 +648,7 @@ py::object te_general_grouped_gemm_for_grouped_tensor( // Ensure that cublasLt handle is created on the correct device, // overriding torch.cuda.set_device calls from user side. // Assumes all tensors passed are on the same device. - at::cuda::CUDAGuard device_guard(workspace_cublas.device()); + CUDAGuard device_guard(workspace_cublas.device()); auto grouped_A = GroupedTensorFromPyTorchGroupedTensor(A); auto grouped_B = GroupedTensorFromPyTorchGroupedTensor(B); @@ -676,7 +676,7 @@ py::object te_general_grouped_gemm_for_grouped_tensor( gemm_config.matmul_config.has_value() ? static_cast(*gemm_config.matmul_config) : nullptr, - at::cuda::getCurrentCUDAStream()); + getCurrentCUDAStream()); }); if (!bias.is_none()) { @@ -685,12 +685,12 @@ py::object te_general_grouped_gemm_for_grouped_tensor( auto te_bias_scale = makeTransformerEngineTensor(*bias_scale); NVTE_SCOPED_GIL_RELEASE({ nvte_grouped_scaled_bias_add(grouped_D.data(), grouped_bias.data(), te_bias_scale.data(), - at::cuda::getCurrentCUDAStream()); + getCurrentCUDAStream()); }); } else { NVTE_SCOPED_GIL_RELEASE({ nvte_grouped_bias_add(grouped_D.data(), grouped_bias.data(), - at::cuda::getCurrentCUDAStream()); + getCurrentCUDAStream()); }); } } @@ -700,10 +700,10 @@ py::object te_general_grouped_gemm_for_grouped_tensor( py::object te_general_grouped_gemm_for_discrete_in(py::handle A, bool transa, py::handle B, bool transb, py::handle D, py::object bias, - std::optional bias_scale, - at::Tensor alpha, at::Tensor beta, - at::Tensor workspace_setup, - at::Tensor workspace_cublas, + std::optional bias_scale, + Tensor alpha, Tensor beta, + Tensor workspace_setup, + Tensor workspace_cublas, bool use_split_accumulator, int math_sm_count) { using namespace transformer_engine::pytorch::detail; @@ -712,7 +712,7 @@ py::object te_general_grouped_gemm_for_discrete_in(py::handle A, bool transa, py // Ensure that cublasLt handle is created on the correct device, // overriding torch.cuda.set_device calls from user side. // Assumes all tensors passed are on the same device. - at::cuda::CUDAGuard device_guard(workspace_cublas.device()); + CUDAGuard device_guard(workspace_cublas.device()); auto grouped_B = GroupedTensorFromPyTorchGroupedTensor(B); auto grouped_D = GroupedTensorFromPyTorchGroupedTensor(D); @@ -738,7 +738,7 @@ py::object te_general_grouped_gemm_for_discrete_in(py::handle A, bool transa, py te_A_vector.emplace_back(te_A_wrappers.back().data()); } - std::vector> swizzled_scale_inverses_list; + std::vector> swizzled_scale_inverses_list; swizzled_scale_inverses_list.emplace_back( multi_tensor_swizzle_scales_for_gemm(te_A_wrappers, transa, !transa)); @@ -753,7 +753,7 @@ py::object te_general_grouped_gemm_for_discrete_in(py::handle A, bool transa, py gemm_config.matmul_config.has_value() ? static_cast(*gemm_config.matmul_config) : nullptr, - at::cuda::getCurrentCUDAStream()); + getCurrentCUDAStream()); }); if (!bias.is_none()) { @@ -762,12 +762,12 @@ py::object te_general_grouped_gemm_for_discrete_in(py::handle A, bool transa, py auto te_bias_scale = makeTransformerEngineTensor(*bias_scale); NVTE_SCOPED_GIL_RELEASE({ nvte_grouped_scaled_bias_add(grouped_D.data(), grouped_bias.data(), te_bias_scale.data(), - at::cuda::getCurrentCUDAStream()); + getCurrentCUDAStream()); }); } else { NVTE_SCOPED_GIL_RELEASE({ nvte_grouped_bias_add(grouped_D.data(), grouped_bias.data(), - at::cuda::getCurrentCUDAStream()); + getCurrentCUDAStream()); }); } } @@ -777,10 +777,10 @@ py::object te_general_grouped_gemm_for_discrete_in(py::handle A, bool transa, py py::object te_general_grouped_gemm_for_discrete_out(py::handle A, bool transa, py::handle B, bool transb, py::handle D, py::object bias, - std::optional bias_scale, - at::Tensor alpha, at::Tensor beta, - at::Tensor workspace_setup, - at::Tensor workspace_cublas, + std::optional bias_scale, + Tensor alpha, Tensor beta, + Tensor workspace_setup, + Tensor workspace_cublas, bool use_split_accumulator, int math_sm_count) { using namespace transformer_engine::pytorch::detail; @@ -789,7 +789,7 @@ py::object te_general_grouped_gemm_for_discrete_out(py::handle A, bool transa, p // Ensure that cublasLt handle is created on the correct device, // overriding torch.cuda.set_device calls from user side. // Assumes all tensors passed are on the same device. - at::cuda::CUDAGuard device_guard(workspace_cublas.device()); + CUDAGuard device_guard(workspace_cublas.device()); NVTE_CHECK(bias.is_none(), "Bias is not supported for discrete output grouped GEMM."); @@ -830,7 +830,7 @@ py::object te_general_grouped_gemm_for_discrete_out(py::handle A, bool transa, p gemm_config.matmul_config.has_value() ? static_cast(*gemm_config.matmul_config) : nullptr, - at::cuda::getCurrentCUDAStream()); + getCurrentCUDAStream()); }); return py::reinterpret_borrow(D); diff --git a/transformer_engine/pytorch/csrc/extensions/grouped_mlp_experimental.cpp b/transformer_engine/pytorch/csrc/extensions/grouped_mlp_experimental.cpp index 0ab8bc8d61..1e805eca9a 100644 --- a/transformer_engine/pytorch/csrc/extensions/grouped_mlp_experimental.cpp +++ b/transformer_engine/pytorch/csrc/extensions/grouped_mlp_experimental.cpp @@ -6,8 +6,6 @@ // Experimental helpers for the fused grouped MLP. -#include - #include #include #include @@ -20,9 +18,9 @@ namespace transformer_engine { namespace pytorch { namespace grouped_mlp_experimental { -std::tuple swizzle_scales_and_pack_ptrs_for_discrete_weights( - const std::vector &data_tensors, const std::vector &scale_tensors, - const std::string &swizzle_type_str, const c10::Device &device) { +std::tuple swizzle_scales_and_pack_ptrs_for_discrete_weights( + const std::vector &data_tensors, const std::vector &scale_tensors, + const std::string &swizzle_type_str, const Device &device) { const size_t num_tensors = data_tensors.size(); NVTE_CHECK(scale_tensors.size() == num_tensors, "Expected data_tensors and scale_tensors to have matching sizes, but got ", @@ -44,13 +42,13 @@ std::tuple swizzle_scales_and_pack_ptrs_for_ // Trivial case: no tensors. Return empty tensors. if (num_tensors == 0) { - auto empty_ptrs = at::empty({0}, at::TensorOptions().dtype(at::kLong).device(device)); - auto empty_scales = at::empty({0}, at::TensorOptions().dtype(at::kByte).device(device)); + auto empty_ptrs = empty({0}, TensorOptions().dtype(kLong).device(device)); + auto empty_scales = empty({0}, TensorOptions().dtype(kByte).device(device)); return {empty_ptrs, empty_ptrs.clone(), std::move(empty_scales)}; } // CUDA stream - auto stream = at::cuda::getCurrentCUDAStream(); + auto stream = getCurrentCUDAStream(); // Tensor properties NVTEScalingMode scaling_mode; @@ -101,8 +99,8 @@ std::tuple swizzle_scales_and_pack_ptrs_for_ // Allocate single buffer for swizzled scales. Uses a uniform stride since // all tensors share the same scale shape. const size_t swizzled_scales_stride = roundup(scale_bytes, 16); // Align to 16 bytes - auto swizzled_scales = at::empty({static_cast(swizzled_scales_stride * num_tensors)}, - at::TensorOptions().dtype(at::kByte).device(device)); + auto swizzled_scales = empty({static_cast(swizzled_scales_stride * num_tensors)}, + TensorOptions().dtype(kByte).device(device)); uint8_t *swizzled_scales_dptr = reinterpret_cast(swizzled_scales.data_ptr()); // Allocate input/output NVTETensors as a single batch. The first @@ -141,8 +139,8 @@ std::tuple swizzle_scales_and_pack_ptrs_for_ packed_ptrs_host[num_tensors + i] = reinterpret_cast(swizzled_scales_dptr + i * swizzled_scales_stride); } - auto packed_ptrs_device = at::empty({static_cast(2 * num_tensors)}, - at::TensorOptions().dtype(at::kLong).device(device)); + auto packed_ptrs_device = empty({static_cast(2 * num_tensors)}, + TensorOptions().dtype(kLong).device(device)); nvte_copy_host_to_device_via_kernel(packed_ptrs_host.data(), packed_ptrs_device.data_ptr(), 2 * num_tensors * sizeof(uint64_t), stream); diff --git a/transformer_engine/pytorch/csrc/extensions/misc.cpp b/transformer_engine/pytorch/csrc/extensions/misc.cpp index ba4371ffe1..97de9e2692 100644 --- a/transformer_engine/pytorch/csrc/extensions/misc.cpp +++ b/transformer_engine/pytorch/csrc/extensions/misc.cpp @@ -4,8 +4,6 @@ * See LICENSE for license information. ************************************************************************/ -#include - #include #include #include @@ -20,27 +18,27 @@ size_t get_cublasLt_version() { return cublasLtGetVersion(); } size_t get_cudnn_version() { return cudnnGetVersion(); } -at::Tensor splits_to_offsets(const at::Tensor &first_dims, int64_t logical_last_dim) { +Tensor splits_to_offsets(const Tensor &first_dims, int64_t logical_last_dim) { NVTE_CHECK(first_dims.is_cuda(), "first_dims must be on CUDA."); - NVTE_CHECK(first_dims.scalar_type() == at::kLong, "first_dims must have dtype int64."); + NVTE_CHECK(first_dims.scalar_type() == kLong, "first_dims must have dtype int64."); NVTE_CHECK(first_dims.dim() == 1, "first_dims must be a 1D tensor."); NVTE_CHECK(logical_last_dim > 0, "logical_last_dim must be greater than 0."); auto first_dims_contiguous = first_dims.contiguous(); const auto num_tensors = static_cast(first_dims_contiguous.numel()); - auto output = at::empty({static_cast(num_tensors) + 1}, - first_dims_contiguous.options().dtype(at::kLong)); + auto output = empty({static_cast(num_tensors) + 1}, + first_dims_contiguous.options().dtype(kLong)); nvte_splits_to_offsets(static_cast(first_dims_contiguous.data_ptr()), static_cast(output.data_ptr()), num_tensors, logical_last_dim, - at::cuda::getCurrentCUDAStream()); + getCurrentCUDAStream()); return output; } -std::tuple> splits_to_offsets_multi( - const at::Tensor &split_sizes, const c10::Device &device, const std::vector &strides, - const std::vector &include_leading_zero, const std::vector &dtypes, +std::tuple> splits_to_offsets_multi( + const Tensor &split_sizes, const Device &device, const std::vector &strides, + const std::vector &include_leading_zero, const std::vector &dtypes, bool bulk_allocate_outputs) { const size_t num_outputs = strides.size(); const size_t num_splits = static_cast(split_sizes.numel()); @@ -52,13 +50,13 @@ std::tuple> splits_to_offsets_multi( NVTE_CHECK(device.is_cuda(), "device must be CUDA, but got ", device.str(), "."); // Convert split sizes to int64 GPU tensor. - const at::Tensor split_sizes_i64 = - split_sizes.scalar_type() == at::kLong ? split_sizes : split_sizes.to(at::kLong); - const at::Tensor split_sizes_out = + const Tensor split_sizes_i64 = + split_sizes.scalar_type() == kLong ? split_sizes : split_sizes.to(kLong); + const Tensor split_sizes_out = split_sizes_i64.device() == device ? split_sizes_i64 : split_sizes_i64.to(device); // Allocate outputs. - std::vector outputs; + std::vector outputs; outputs.reserve(num_outputs); if (bulk_allocate_outputs) { std::vector> shapes; @@ -75,7 +73,7 @@ std::tuple> splits_to_offsets_multi( for (size_t i = 0; i < num_outputs; ++i) { const int64_t length = static_cast(num_splits) + (include_leading_zero[i] ? 1 : 0); outputs.emplace_back( - at::empty({length}, at::TensorOptions().dtype(dtypes[i]).device(device))); + empty({length}, TensorOptions().dtype(dtypes[i]).device(device))); } } @@ -95,14 +93,14 @@ std::tuple> splits_to_offsets_multi( NVTE_SCOPED_GIL_RELEASE({ nvte_splits_to_offsets_multi(split_sizes_nvte.data(), outputs_nvte.data(), strides.data(), include_leading_zero_int.data(), num_outputs, - at::cuda::getCurrentCUDAStream()); + getCurrentCUDAStream()); }); return {split_sizes_out, std::move(outputs)}; } -at::Tensor copy_data_ptrs_to_device(const std::vector &tensors, - const c10::Device &device) { +Tensor copy_data_ptrs_to_device(const std::vector &tensors, + const Device &device) { // Collect data pointers std::vector ptrs_host; ptrs_host.reserve(tensors.size()); @@ -111,13 +109,13 @@ at::Tensor copy_data_ptrs_to_device(const std::vector &tensors, } // Allocate device buffer - auto ptrs_device = at::empty({static_cast(tensors.size())}, - at::TensorOptions().dtype(at::kLong).device(device)); + auto ptrs_device = empty({static_cast(tensors.size())}, + TensorOptions().dtype(kLong).device(device)); // Load pointers on device nvte_copy_host_to_device_via_kernel(ptrs_host.data(), ptrs_device.data_ptr(), tensors.size() * sizeof(uint64_t), - at::cuda::getCurrentCUDAStream()); + getCurrentCUDAStream()); return ptrs_device; } diff --git a/transformer_engine/pytorch/csrc/extensions/multi_tensor/adam.cpp b/transformer_engine/pytorch/csrc/extensions/multi_tensor/adam.cpp index 145e1d4b40..a790936e7c 100644 --- a/transformer_engine/pytorch/csrc/extensions/multi_tensor/adam.cpp +++ b/transformer_engine/pytorch/csrc/extensions/multi_tensor/adam.cpp @@ -8,8 +8,8 @@ namespace transformer_engine::pytorch { -void multi_tensor_adam_cuda(int chunk_size, at::Tensor noop_flag, - std::vector> tensor_lists, const float lr, +void multi_tensor_adam_cuda(int chunk_size, Tensor noop_flag, + std::vector> tensor_lists, const float lr, const float beta1, const float beta2, const float epsilon, const int step, const int mode, const int bias_correction, const float weight_decay) { @@ -19,11 +19,11 @@ void multi_tensor_adam_cuda(int chunk_size, at::Tensor noop_flag, nvte_multi_tensor_adam_cuda(chunk_size, noop_flag_cu.data(), tensor_lists_ptr.data(), num_lists, num_tensors, lr, beta1, beta2, epsilon, step, mode, bias_correction, - weight_decay, at::cuda::getCurrentCUDAStream()); + weight_decay, getCurrentCUDAStream()); } -void multi_tensor_adam_param_remainder_cuda(int chunk_size, at::Tensor noop_flag, - std::vector> tensor_lists, +void multi_tensor_adam_param_remainder_cuda(int chunk_size, Tensor noop_flag, + std::vector> tensor_lists, const float lr, const float beta1, const float beta2, const float epsilon, const int step, const int mode, const int bias_correction, const float weight_decay) { @@ -33,11 +33,11 @@ void multi_tensor_adam_param_remainder_cuda(int chunk_size, at::Tensor noop_flag nvte_multi_tensor_adam_param_remainder_cuda( chunk_size, noop_flag_cu.data(), tensor_lists_ptr.data(), num_lists, num_tensors, lr, beta1, - beta2, epsilon, step, mode, bias_correction, weight_decay, at::cuda::getCurrentCUDAStream()); + beta2, epsilon, step, mode, bias_correction, weight_decay, getCurrentCUDAStream()); } -void multi_tensor_adam_fp8_cuda(int chunk_size, at::Tensor noop_flag, - std::vector> tensor_lists, const float lr, +void multi_tensor_adam_fp8_cuda(int chunk_size, Tensor noop_flag, + std::vector> tensor_lists, const float lr, const float beta1, const float beta2, const float epsilon, const int step, const int mode, const int bias_correction, const float weight_decay, DType fp8_dtype) { @@ -48,15 +48,15 @@ void multi_tensor_adam_fp8_cuda(int chunk_size, at::Tensor noop_flag, nvte_multi_tensor_adam_fp8_cuda(chunk_size, noop_flag_cu.data(), tensor_lists_ptr.data(), num_lists, num_tensors, lr, beta1, beta2, epsilon, step, mode, bias_correction, weight_decay, static_cast(fp8_dtype), - at::cuda::getCurrentCUDAStream()); + getCurrentCUDAStream()); } -void multi_tensor_adam_capturable_cuda(int chunk_size, at::Tensor noop_flag, - std::vector> tensor_lists, - at::Tensor lr, const float beta1, const float beta2, - const float epsilon, at::Tensor step, const int mode, +void multi_tensor_adam_capturable_cuda(int chunk_size, Tensor noop_flag, + std::vector> tensor_lists, + Tensor lr, const float beta1, const float beta2, + const float epsilon, Tensor step, const int mode, const int bias_correction, const float weight_decay, - at::Tensor inv_scale) { + Tensor inv_scale) { auto noop_flag_cu = makeTransformerEngineTensor(noop_flag); auto [_, __, tensor_lists_ptr, num_lists, num_tensors] = makeTransformerEngineTensorList(tensor_lists); @@ -67,15 +67,15 @@ void multi_tensor_adam_capturable_cuda(int chunk_size, at::Tensor noop_flag, nvte_multi_tensor_adam_capturable_cuda( chunk_size, noop_flag_cu.data(), tensor_lists_ptr.data(), num_lists, num_tensors, lr_cu.data(), beta1, beta2, epsilon, step_cu.data(), mode, bias_correction, weight_decay, - inv_scale_cu.data(), at::cuda::getCurrentCUDAStream()); + inv_scale_cu.data(), getCurrentCUDAStream()); } -void multi_tensor_adam_capturable_master_cuda(int chunk_size, at::Tensor noop_flag, - std::vector> tensor_lists, - at::Tensor lr, const float beta1, const float beta2, - const float epsilon, at::Tensor step, const int mode, +void multi_tensor_adam_capturable_master_cuda(int chunk_size, Tensor noop_flag, + std::vector> tensor_lists, + Tensor lr, const float beta1, const float beta2, + const float epsilon, Tensor step, const int mode, const int bias_correction, const float weight_decay, - at::Tensor inv_scale) { + Tensor inv_scale) { auto noop_flag_cu = makeTransformerEngineTensor(noop_flag); auto [_, __, tensor_lists_ptr, num_lists, num_tensors] = makeTransformerEngineTensorList(tensor_lists); @@ -86,7 +86,7 @@ void multi_tensor_adam_capturable_master_cuda(int chunk_size, at::Tensor noop_fl nvte_multi_tensor_adam_capturable_master_cuda( chunk_size, noop_flag_cu.data(), tensor_lists_ptr.data(), num_lists, num_tensors, lr_cu.data(), beta1, beta2, epsilon, step_cu.data(), mode, bias_correction, weight_decay, - inv_scale_cu.data(), at::cuda::getCurrentCUDAStream()); + inv_scale_cu.data(), getCurrentCUDAStream()); } } // namespace transformer_engine::pytorch diff --git a/transformer_engine/pytorch/csrc/extensions/multi_tensor/compute_scale.cpp b/transformer_engine/pytorch/csrc/extensions/multi_tensor/compute_scale.cpp index 328970ffa8..257cc8024e 100644 --- a/transformer_engine/pytorch/csrc/extensions/multi_tensor/compute_scale.cpp +++ b/transformer_engine/pytorch/csrc/extensions/multi_tensor/compute_scale.cpp @@ -9,7 +9,7 @@ namespace transformer_engine::pytorch { void multi_tensor_compute_scale_and_scale_inv_cuda( - int chunk_size, at::Tensor noop_flag, std::vector> tensor_lists, + int chunk_size, Tensor noop_flag, std::vector> tensor_lists, float max_fp8, bool force_pow_2_scales, float epsilon) { auto noop_flag_cu = makeTransformerEngineTensor(noop_flag); auto [_, __, tensor_lists_ptr, num_lists, num_tensors] = @@ -17,17 +17,17 @@ void multi_tensor_compute_scale_and_scale_inv_cuda( nvte_multi_tensor_compute_scale_and_scale_inv_cuda( chunk_size, noop_flag_cu.data(), tensor_lists_ptr.data(), num_lists, num_tensors, max_fp8, - force_pow_2_scales, epsilon, at::cuda::getCurrentCUDAStream()); + force_pow_2_scales, epsilon, getCurrentCUDAStream()); } void multi_tensor_compute_scale_inv_e8m0_cuda(int chunk_size, const py::object &dummy, - std::vector> tensor_lists) { + std::vector> tensor_lists) { NVTE_CHECK(dummy.is_none(), "No-op flag is not supported."); auto [_, __, tensor_lists_ptr, num_lists, num_tensors] = makeTransformerEngineTensorList(tensor_lists); nvte_multi_tensor_compute_scale_inv_e8m0_cuda(chunk_size, tensor_lists_ptr.data(), num_lists, - num_tensors, at::cuda::getCurrentCUDAStream()); + num_tensors, getCurrentCUDAStream()); } } // namespace transformer_engine::pytorch diff --git a/transformer_engine/pytorch/csrc/extensions/multi_tensor/l2norm.cpp b/transformer_engine/pytorch/csrc/extensions/multi_tensor/l2norm.cpp index b02cf1fbba..c4c6907a2d 100644 --- a/transformer_engine/pytorch/csrc/extensions/multi_tensor/l2norm.cpp +++ b/transformer_engine/pytorch/csrc/extensions/multi_tensor/l2norm.cpp @@ -8,17 +8,17 @@ namespace transformer_engine::pytorch { -std::tuple multi_tensor_l2norm_cuda( - int chunk_size, at::Tensor noop_flag, std::vector> tensor_lists, - at::optional per_tensor_python) { +std::tuple multi_tensor_l2norm_cuda( + int chunk_size, Tensor noop_flag, std::vector> tensor_lists, + std::optional per_tensor_python) { bool per_tensor = per_tensor_python.has_value() ? per_tensor_python.value() : false; - auto float_options = tensor_lists[0][0].options().dtype(at::kFloat); - auto output = at::zeros({320}, float_options); + auto float_options = tensor_lists[0][0].options().dtype(kFloat); + auto output = zeros({320}, float_options); - at::Tensor output_per_tensor; - at::Tensor ret_per_tensor; - auto ret = at::empty({1}, output.options()); + Tensor output_per_tensor; + Tensor ret_per_tensor; + auto ret = empty({1}, output.options()); int ntensors = tensor_lists[0].size(); int max_chunks_per_tensor = -1; @@ -29,11 +29,11 @@ std::tuple multi_tensor_l2norm_cuda( if (max_chunks_this_tensor > max_chunks_per_tensor) max_chunks_per_tensor = max_chunks_this_tensor; } - output_per_tensor = at::zeros({ntensors * max_chunks_per_tensor}, float_options); - ret_per_tensor = at::empty({ntensors}, float_options); + output_per_tensor = zeros({ntensors * max_chunks_per_tensor}, float_options); + ret_per_tensor = empty({ntensors}, float_options); } else { - output_per_tensor = at::empty({0}, float_options); - ret_per_tensor = at::empty({0}, float_options); + output_per_tensor = empty({0}, float_options); + ret_per_tensor = empty({0}, float_options); } auto noop_flag_cu = makeTransformerEngineTensor(noop_flag); @@ -47,21 +47,21 @@ std::tuple multi_tensor_l2norm_cuda( nvte_multi_tensor_l2norm_cuda(chunk_size, noop_flag_cu.data(), tensor_lists_ptr.data(), num_lists, num_tensors, output_cu.data(), output_per_tensor_cu.data(), ret_cu.data(), ret_per_tensor_cu.data(), per_tensor, - max_chunks_per_tensor, at::cuda::getCurrentCUDAStream()); + max_chunks_per_tensor, getCurrentCUDAStream()); - return std::tuple(ret, ret_per_tensor); + return std::tuple(ret, ret_per_tensor); } -std::tuple multi_tensor_unscale_l2norm_cuda( - int chunk_size, at::Tensor noop_flag, std::vector> tensor_lists, - at::Tensor inv_scale, at::optional per_tensor_python) { +std::tuple multi_tensor_unscale_l2norm_cuda( + int chunk_size, Tensor noop_flag, std::vector> tensor_lists, + Tensor inv_scale, std::optional per_tensor_python) { bool per_tensor = per_tensor_python.has_value() ? per_tensor_python.value() : false; - auto float_options = tensor_lists[0][0].options().dtype(at::kFloat); - auto output = at::zeros({320}, float_options); + auto float_options = tensor_lists[0][0].options().dtype(kFloat); + auto output = zeros({320}, float_options); - at::Tensor output_per_tensor; - at::Tensor ret_per_tensor; + Tensor output_per_tensor; + Tensor ret_per_tensor; int ntensors = tensor_lists[0].size(); int max_chunks_per_tensor = -1; @@ -73,14 +73,14 @@ std::tuple multi_tensor_unscale_l2norm_cuda( if (max_chunks_this_tensor > max_chunks_per_tensor) max_chunks_per_tensor = max_chunks_this_tensor; } - output_per_tensor = at::zeros({ntensors * max_chunks_per_tensor}, float_options); - ret_per_tensor = at::empty({ntensors}, float_options); + output_per_tensor = zeros({ntensors * max_chunks_per_tensor}, float_options); + ret_per_tensor = empty({ntensors}, float_options); } else { - output_per_tensor = at::empty({0}, float_options); - ret_per_tensor = at::empty({0}, float_options); + output_per_tensor = empty({0}, float_options); + ret_per_tensor = empty({0}, float_options); } - auto ret = at::empty({1}, output.options()); + auto ret = empty({1}, output.options()); auto noop_flag_cu = makeTransformerEngineTensor(noop_flag); auto [_, __, tensor_lists_ptr, num_lists, num_tensors] = @@ -94,9 +94,9 @@ std::tuple multi_tensor_unscale_l2norm_cuda( nvte_multi_tensor_unscale_l2norm_cuda( chunk_size, noop_flag_cu.data(), tensor_lists_ptr.data(), num_lists, num_tensors, output_cu.data(), output_per_tensor_cu.data(), ret_cu.data(), ret_per_tensor_cu.data(), - inv_scale_cu.data(), per_tensor, max_chunks_per_tensor, at::cuda::getCurrentCUDAStream()); + inv_scale_cu.data(), per_tensor, max_chunks_per_tensor, getCurrentCUDAStream()); - return std::tuple(ret, ret_per_tensor); + return std::tuple(ret, ret_per_tensor); } } // namespace transformer_engine::pytorch diff --git a/transformer_engine/pytorch/csrc/extensions/multi_tensor/scale.cpp b/transformer_engine/pytorch/csrc/extensions/multi_tensor/scale.cpp index 687eb34f32..bc402eca75 100644 --- a/transformer_engine/pytorch/csrc/extensions/multi_tensor/scale.cpp +++ b/transformer_engine/pytorch/csrc/extensions/multi_tensor/scale.cpp @@ -8,26 +8,26 @@ namespace transformer_engine::pytorch { -void multi_tensor_scale_cuda(int chunk_size, at::Tensor is_infinite, - std::vector> tensor_lists, float scale) { +void multi_tensor_scale_cuda(int chunk_size, Tensor is_infinite, + std::vector> tensor_lists, float scale) { auto is_infinite_cu = makeTransformerEngineTensor(is_infinite); auto [_, __, tensor_lists_ptr, num_lists, num_tensors] = makeTransformerEngineTensorList(tensor_lists); nvte_multi_tensor_scale_cuda(chunk_size, is_infinite_cu.data(), tensor_lists_ptr.data(), - num_lists, num_tensors, scale, at::cuda::getCurrentCUDAStream()); + num_lists, num_tensors, scale, getCurrentCUDAStream()); } -void multi_tensor_scale_tensor_cuda(int chunk_size, at::Tensor is_infinite, - std::vector> tensor_lists, - at::Tensor scale) { +void multi_tensor_scale_tensor_cuda(int chunk_size, Tensor is_infinite, + std::vector> tensor_lists, + Tensor scale) { auto is_infinite_cu = makeTransformerEngineTensor(is_infinite); auto scale_cu = makeTransformerEngineTensor(scale); auto [_, __, tensor_lists_ptr, num_lists, num_tensors] = makeTransformerEngineTensorList(tensor_lists); nvte_multi_tensor_scale_tensor_cuda(chunk_size, is_infinite_cu.data(), tensor_lists_ptr.data(), num_lists, num_tensors, scale_cu.data(), - at::cuda::getCurrentCUDAStream()); + getCurrentCUDAStream()); } } // namespace transformer_engine::pytorch diff --git a/transformer_engine/pytorch/csrc/extensions/multi_tensor/sgd.cpp b/transformer_engine/pytorch/csrc/extensions/multi_tensor/sgd.cpp index a70fe12b56..4dbbb4124d 100644 --- a/transformer_engine/pytorch/csrc/extensions/multi_tensor/sgd.cpp +++ b/transformer_engine/pytorch/csrc/extensions/multi_tensor/sgd.cpp @@ -8,8 +8,8 @@ namespace transformer_engine::pytorch { -void multi_tensor_sgd_cuda(int chunk_size, at::Tensor noop_flag, - std::vector> tensor_lists, float wd, +void multi_tensor_sgd_cuda(int chunk_size, Tensor noop_flag, + std::vector> tensor_lists, float wd, float momentum, float dampening, float lr, bool nesterov, bool first_run, bool wd_after_momentum, float scale) { auto noop_flag_cu = makeTransformerEngineTensor(noop_flag); @@ -18,7 +18,7 @@ void multi_tensor_sgd_cuda(int chunk_size, at::Tensor noop_flag, nvte_multi_tensor_sgd_cuda(chunk_size, noop_flag_cu.data(), tensor_lists_ptr.data(), num_lists, num_tensors, wd, momentum, dampening, lr, nesterov, first_run, - wd_after_momentum, scale, at::cuda::getCurrentCUDAStream()); + wd_after_momentum, scale, getCurrentCUDAStream()); } } // namespace transformer_engine::pytorch diff --git a/transformer_engine/pytorch/csrc/extensions/newton_schulz.cpp b/transformer_engine/pytorch/csrc/extensions/newton_schulz.cpp index 8b24e8fdb9..26c0b25b54 100644 --- a/transformer_engine/pytorch/csrc/extensions/newton_schulz.cpp +++ b/transformer_engine/pytorch/csrc/extensions/newton_schulz.cpp @@ -21,7 +21,7 @@ void cusolvermp_ctx_destroy(int64_t ctx_ptr) { nvte_cusolvermp_ctx_destroy(ctx); } -void newton_schulz(int64_t ctx_ptr, int64_t m, int64_t n, at::Tensor x, int64_t num_iterations, +void newton_schulz(int64_t ctx_ptr, int64_t m, int64_t n, Tensor x, int64_t num_iterations, std::vector coefficients) { auto* ctx = reinterpret_cast(ctx_ptr); @@ -32,7 +32,7 @@ void newton_schulz(int64_t ctx_ptr, int64_t m, int64_t n, at::Tensor x, int64_t auto te_dtype = GetTransformerEngineDType(x.scalar_type()); TensorWrapper x_tensor(x.data_ptr(), shape, te_dtype); - auto caller_stream = at::cuda::getCurrentCUDAStream().stream(); + auto caller_stream = getCurrentCUDAStream().stream(); nvte_newton_schulz(ctx, m, n, x_tensor.data(), num_iterations, coefficients.data(), static_cast(coefficients.size()), caller_stream); } diff --git a/transformer_engine/pytorch/csrc/extensions/normalization.cpp b/transformer_engine/pytorch/csrc/extensions/normalization.cpp index c3dec944e4..d8d7e69c03 100644 --- a/transformer_engine/pytorch/csrc/extensions/normalization.cpp +++ b/transformer_engine/pytorch/csrc/extensions/normalization.cpp @@ -10,9 +10,9 @@ namespace transformer_engine::pytorch { -std::vector layernorm_bwd(const at::Tensor &dz, const at::Tensor &x, - const at::Tensor &mu, const at::Tensor &rsigma, - const at::Tensor &gamma, const int sm_margin, +std::vector layernorm_bwd(const Tensor &dz, const Tensor &x, + const Tensor &mu, const Tensor &rsigma, + const Tensor &gamma, const int sm_margin, const bool zero_centered_gamma) { const auto &dz_ = dz.contiguous(); const auto &x_ = x.contiguous(); @@ -20,9 +20,9 @@ std::vector layernorm_bwd(const at::Tensor &dz, const at::Tensor &x, const auto &rsigma_ = rsigma.contiguous(); const auto &gamma_ = gamma.contiguous(); - auto dx = at::empty_like(x_); - auto dgamma = at::empty_like(gamma_); - auto dbeta = at::empty_like(gamma_); + auto dx = empty_like(x_); + auto dgamma = empty_like(gamma_); + auto dbeta = empty_like(gamma_); TensorWrapper workspace; auto dz_cu = makeTransformerEngineTensor(dz_); @@ -38,8 +38,8 @@ std::vector layernorm_bwd(const at::Tensor &dz, const at::Tensor &x, NVTE_SCOPED_GIL_RELEASE({ nvte_layernorm_bwd(dz_cu.data(), x_cu.data(), mu_cu.data(), rsigma_cu.data(), gamma_cu.data(), dx_cu.data(), dgamma_cu.data(), dbeta_cu.data(), workspace.data(), - at::cuda::getCurrentDeviceProperties()->multiProcessorCount - sm_margin, - zero_centered_gamma, at::cuda::getCurrentCUDAStream()); + getCurrentDeviceProperties()->multiProcessorCount - sm_margin, + zero_centered_gamma, getCurrentCUDAStream()); }); // Alloc space for Tensors. @@ -51,8 +51,8 @@ std::vector layernorm_bwd(const at::Tensor &dz, const at::Tensor &x, NVTE_SCOPED_GIL_RELEASE({ nvte_layernorm_bwd(dz_cu.data(), x_cu.data(), mu_cu.data(), rsigma_cu.data(), gamma_cu.data(), dx_cu.data(), dgamma_cu.data(), dbeta_cu.data(), workspace.data(), - at::cuda::getCurrentDeviceProperties()->multiProcessorCount - sm_margin, - zero_centered_gamma, at::cuda::getCurrentCUDAStream()); + getCurrentDeviceProperties()->multiProcessorCount - sm_margin, + zero_centered_gamma, getCurrentCUDAStream()); }); return {py::cast(dx), py::cast(dgamma), py::cast(dbeta)}; @@ -67,7 +67,7 @@ std::vector layernorm_fwd(py::handle input, py::handle weight, Maybe // Ensure that cuDNN handle is created on the correct device, // overriding torch.cuda.set_device calls from user side. // Assumes all tensors passed are on the same device. - at::cuda::CUDAGuard device_guard(input.cast().device()); + CUDAGuard device_guard(input.cast().device()); // Input and param tensors auto none = py::none(); @@ -83,8 +83,8 @@ std::vector layernorm_fwd(py::handle input, py::handle weight, Maybe const auto [outer_size, inner_size] = get_2d_dims(shape); // Tensors to save for backward pass - at::Tensor mu_py = at::empty({static_cast(outer_size)}, at::CUDA(at::kFloat)); - at::Tensor rsigma_py = at::empty({static_cast(outer_size)}, at::CUDA(at::kFloat)); + Tensor mu_py = empty({static_cast(outer_size)}, CUDA(kFloat)); + Tensor rsigma_py = empty({static_cast(outer_size)}, CUDA(kFloat)); TensorWrapper mu_nvte = makeTransformerEngineTensor(mu_py); TensorWrapper rsigma_nvte = makeTransformerEngineTensor(rsigma_py); @@ -145,7 +145,7 @@ std::vector layernorm_fwd(py::handle input, py::handle weight, Maybe // Construct unquantized output tensor if needed TensorWrapper unquantized_out_nvte; py::object unquantized_out; - at::Tensor amax_buf; + Tensor amax_buf; TensorWrapper *kernel_out_nvte = &out_nvte; switch (impl) { case Impl::UNFUSED: { @@ -175,8 +175,8 @@ std::vector layernorm_fwd(py::handle input, py::handle weight, Maybe nvte_layernorm_fwd(input_nvte.data(), weight_nvte.data(), bias_nvte.data(), eps, kernel_out_nvte->data(), mu_nvte.data(), rsigma_nvte.data(), workspace.data(), - at::cuda::getCurrentDeviceProperties()->multiProcessorCount - sm_margin, - zero_centered_gamma, at::cuda::getCurrentCUDAStream()); + getCurrentDeviceProperties()->multiProcessorCount - sm_margin, + zero_centered_gamma, getCurrentCUDAStream()); }); // Allocate workspace @@ -189,8 +189,8 @@ std::vector layernorm_fwd(py::handle input, py::handle weight, Maybe nvte_layernorm_fwd(input_nvte.data(), weight_nvte.data(), bias_nvte.data(), eps, kernel_out_nvte->data(), mu_nvte.data(), rsigma_nvte.data(), workspace.data(), - at::cuda::getCurrentDeviceProperties()->multiProcessorCount - sm_margin, - zero_centered_gamma, at::cuda::getCurrentCUDAStream()); + getCurrentDeviceProperties()->multiProcessorCount - sm_margin, + zero_centered_gamma, getCurrentCUDAStream()); }); // Quantize output if needed @@ -213,16 +213,16 @@ std::vector layernorm_fwd(py::handle input, py::handle weight, Maybe return {out, py::cast(mu_py), py::cast(rsigma_py)}; } -std::vector rmsnorm_bwd(const at::Tensor &dz, const at::Tensor &x, - const at::Tensor &rsigma, const at::Tensor &gamma, +std::vector rmsnorm_bwd(const Tensor &dz, const Tensor &x, + const Tensor &rsigma, const Tensor &gamma, const int sm_margin, const bool zero_centered_gamma) { const auto &dz_ = dz.contiguous(); const auto &x_ = x.contiguous(); const auto &rsigma_ = rsigma.contiguous(); const auto &gamma_ = gamma.contiguous(); - auto dx = at::empty_like(x_); - auto dgamma = at::empty_like(gamma_); + auto dx = empty_like(x_); + auto dgamma = empty_like(gamma_); TensorWrapper workspace; auto dz_cu = makeTransformerEngineTensor(dz_); @@ -236,8 +236,8 @@ std::vector rmsnorm_bwd(const at::Tensor &dz, const at::Tensor &x, NVTE_SCOPED_GIL_RELEASE({ nvte_rmsnorm_bwd(dz_cu.data(), x_cu.data(), rsigma_cu.data(), gamma_cu.data(), dx_cu.data(), dgamma_cu.data(), workspace.data(), - at::cuda::getCurrentDeviceProperties()->multiProcessorCount - sm_margin, - zero_centered_gamma, at::cuda::getCurrentCUDAStream()); + getCurrentDeviceProperties()->multiProcessorCount - sm_margin, + zero_centered_gamma, getCurrentCUDAStream()); }); // Alloc space for Tensors. @@ -249,16 +249,16 @@ std::vector rmsnorm_bwd(const at::Tensor &dz, const at::Tensor &x, NVTE_SCOPED_GIL_RELEASE({ nvte_rmsnorm_bwd(dz_cu.data(), x_cu.data(), rsigma_cu.data(), gamma_cu.data(), dx_cu.data(), dgamma_cu.data(), workspace.data(), - at::cuda::getCurrentDeviceProperties()->multiProcessorCount - sm_margin, - zero_centered_gamma, at::cuda::getCurrentCUDAStream()); + getCurrentDeviceProperties()->multiProcessorCount - sm_margin, + zero_centered_gamma, getCurrentCUDAStream()); }); return {py::cast(dx), py::cast(dgamma)}; } -std::vector rmsnorm_bwd_add(const at::Tensor &dz, const at::Tensor &x, - const at::Tensor &add, const at::Tensor &rsigma, - const at::Tensor &gamma, const int sm_margin, +std::vector rmsnorm_bwd_add(const Tensor &dz, const Tensor &x, + const Tensor &add, const Tensor &rsigma, + const Tensor &gamma, const int sm_margin, const bool zero_centered_gamma) { const auto &dz_ = dz.contiguous(); const auto &x_ = x.contiguous(); @@ -266,8 +266,8 @@ std::vector rmsnorm_bwd_add(const at::Tensor &dz, const at::Tensor & const auto &rsigma_ = rsigma.contiguous(); const auto &gamma_ = gamma.contiguous(); - auto dx = at::empty_like(x_); - auto dgamma = at::empty_like(gamma_); + auto dx = empty_like(x_); + auto dgamma = empty_like(gamma_); TensorWrapper workspace; auto dz_cu = makeTransformerEngineTensor(dz_); @@ -282,8 +282,8 @@ std::vector rmsnorm_bwd_add(const at::Tensor &dz, const at::Tensor & NVTE_SCOPED_GIL_RELEASE({ nvte_rmsnorm_bwd_add(dz_cu.data(), x_cu.data(), add_cu.data(), rsigma_cu.data(), gamma_cu.data(), dx_cu.data(), dgamma_cu.data(), workspace.data(), - at::cuda::getCurrentDeviceProperties()->multiProcessorCount - sm_margin, - zero_centered_gamma, at::cuda::getCurrentCUDAStream()); + getCurrentDeviceProperties()->multiProcessorCount - sm_margin, + zero_centered_gamma, getCurrentCUDAStream()); }); // Alloc space for Tensors. @@ -295,8 +295,8 @@ std::vector rmsnorm_bwd_add(const at::Tensor &dz, const at::Tensor & NVTE_SCOPED_GIL_RELEASE({ nvte_rmsnorm_bwd_add(dz_cu.data(), x_cu.data(), add_cu.data(), rsigma_cu.data(), gamma_cu.data(), dx_cu.data(), dgamma_cu.data(), workspace.data(), - at::cuda::getCurrentDeviceProperties()->multiProcessorCount - sm_margin, - zero_centered_gamma, at::cuda::getCurrentCUDAStream()); + getCurrentDeviceProperties()->multiProcessorCount - sm_margin, + zero_centered_gamma, getCurrentCUDAStream()); }); return {py::cast(dx), py::cast(dgamma)}; @@ -310,7 +310,7 @@ std::vector rmsnorm_fwd(const py::handle &input, const py::handle &w // Ensure that cuDNN handle is created on the correct device, // overriding torch.cuda.set_device calls from user side. // Assumes all tensors passed are on the same device. - at::cuda::CUDAGuard device_guard(input.cast().device()); + CUDAGuard device_guard(input.cast().device()); // Input and param tensors auto none = py::none(); @@ -322,7 +322,7 @@ std::vector rmsnorm_fwd(const py::handle &input, const py::handle &w const auto [outer_size, inner_size] = get_2d_dims(shape); // Tensors to save for backward pass - at::Tensor rsigma_py = at::empty({static_cast(outer_size)}, at::CUDA(at::kFloat)); + Tensor rsigma_py = empty({static_cast(outer_size)}, CUDA(kFloat)); TensorWrapper rsigma_nvte = makeTransformerEngineTensor(rsigma_py); // Quantizer @@ -382,7 +382,7 @@ std::vector rmsnorm_fwd(const py::handle &input, const py::handle &w // Construct unquantized output tensor if needed TensorWrapper unquantized_out_nvte; py::object unquantized_out; - at::Tensor amax_buf; + Tensor amax_buf; TensorWrapper *kernel_out_nvte = &out_nvte; switch (impl) { case Impl::UNFUSED: { @@ -411,8 +411,8 @@ std::vector rmsnorm_fwd(const py::handle &input, const py::handle &w NVTE_SCOPED_GIL_RELEASE({ nvte_rmsnorm_fwd(input_nvte.data(), weight_nvte.data(), eps, kernel_out_nvte->data(), rsigma_nvte.data(), workspace.data(), - at::cuda::getCurrentDeviceProperties()->multiProcessorCount - sm_margin, - zero_centered_gamma, at::cuda::getCurrentCUDAStream()); + getCurrentDeviceProperties()->multiProcessorCount - sm_margin, + zero_centered_gamma, getCurrentCUDAStream()); }); // Allocate workspace @@ -424,8 +424,8 @@ std::vector rmsnorm_fwd(const py::handle &input, const py::handle &w NVTE_SCOPED_GIL_RELEASE({ nvte_rmsnorm_fwd(input_nvte.data(), weight_nvte.data(), eps, kernel_out_nvte->data(), rsigma_nvte.data(), workspace.data(), - at::cuda::getCurrentDeviceProperties()->multiProcessorCount - sm_margin, - zero_centered_gamma, at::cuda::getCurrentCUDAStream()); + getCurrentDeviceProperties()->multiProcessorCount - sm_margin, + zero_centered_gamma, getCurrentCUDAStream()); }); // Quantize output if needed diff --git a/transformer_engine/pytorch/csrc/extensions/nvfp4_2d_partial_cast.cpp b/transformer_engine/pytorch/csrc/extensions/nvfp4_2d_partial_cast.cpp index 685250d137..bfa7bda19e 100644 --- a/transformer_engine/pytorch/csrc/extensions/nvfp4_2d_partial_cast.cpp +++ b/transformer_engine/pytorch/csrc/extensions/nvfp4_2d_partial_cast.cpp @@ -8,13 +8,13 @@ namespace transformer_engine::pytorch { -void nvfp4_2d_compute_partial_amax(const at::Tensor& tensor, at::Tensor amax, size_t h, size_t w, +void nvfp4_2d_compute_partial_amax(const Tensor& tensor, Tensor amax, size_t h, size_t w, size_t start_offset, size_t block_len) { TORCH_CHECK(block_len == 16, "Currently only block_len = 16 is supported for NVFP4 2D"); TORCH_CHECK(amax.dim() == 2, "amax must be a 2D tensor"); - TORCH_CHECK(amax.scalar_type() == at::ScalarType::Float, "amax must be a float tensor"); - TORCH_CHECK(tensor.scalar_type() == at::ScalarType::Float || - tensor.scalar_type() == at::ScalarType::BFloat16, + TORCH_CHECK(amax.scalar_type() == ScalarType::Float, "amax must be a float tensor"); + TORCH_CHECK(tensor.scalar_type() == ScalarType::Float || + tensor.scalar_type() == ScalarType::BFloat16, "tensor must be a float or bfloat16 tensor"); const TensorWrapper tensor_cu = makeTransformerEngineTensor(tensor.contiguous()); @@ -22,20 +22,20 @@ void nvfp4_2d_compute_partial_amax(const at::Tensor& tensor, at::Tensor amax, si nvte_nvfp4_2d_compute_partial_amax(tensor_cu.data(), amax_cu.data(), h, w, amax.stride(0), amax.stride(1), start_offset, block_len, - at::cuda::getCurrentCUDAStream()); + getCurrentCUDAStream()); } -void nvfp4_2d_partial_cast(const at::Tensor& inp, py::handle out, const at::Tensor& scale, - const at::Tensor& global_scale, size_t h, size_t w, size_t start_offset, +void nvfp4_2d_partial_cast(const Tensor& inp, py::handle out, const Tensor& scale, + const Tensor& global_scale, size_t h, size_t w, size_t start_offset, size_t block_len) { TORCH_CHECK(block_len == 16, "Currently only block_len = 16 is supported for NVFP4 2D"); TORCH_CHECK(scale.dim() == 2, "scale must be a 2D tensor"); - TORCH_CHECK(scale.scalar_type() == at::ScalarType::Float, "scale must be a float tensor"); + TORCH_CHECK(scale.scalar_type() == ScalarType::Float, "scale must be a float tensor"); TORCH_CHECK(global_scale.numel() == 1, "global_scale must be a scalar tensor"); - TORCH_CHECK(global_scale.scalar_type() == at::ScalarType::Float, + TORCH_CHECK(global_scale.scalar_type() == ScalarType::Float, "global_scale must be a float tensor"); TORCH_CHECK( - inp.scalar_type() == at::ScalarType::Float || inp.scalar_type() == at::ScalarType::BFloat16, + inp.scalar_type() == ScalarType::Float || inp.scalar_type() == ScalarType::BFloat16, "input must be a float or bfloat16 tensor"); const TensorWrapper inp_cu = makeTransformerEngineTensor(inp.contiguous()); @@ -45,13 +45,13 @@ void nvfp4_2d_partial_cast(const at::Tensor& inp, py::handle out, const at::Tens nvte_nvfp4_2d_partial_cast(inp_cu.data(), out_cu.data(), scale_cu.data(), global_scale_cu.data(), h, w, scale.stride(0), scale.stride(1), start_offset, block_len, - at::cuda::getCurrentCUDAStream()); + getCurrentCUDAStream()); } -void nvfp4_multi_tensor_2d_partial_cast(std::vector inp_list, - std::vector out_list, - std::vector scale_list, - std::vector global_scale_list, +void nvfp4_multi_tensor_2d_partial_cast(std::vector inp_list, + std::vector out_list, + std::vector scale_list, + std::vector global_scale_list, std::vector h_list, std::vector w_list, std::vector start_offset_list, int64_t block_len) { TORCH_CHECK(block_len == 16, "Currently only block_len = 16 is supported for NVFP4 2D"); @@ -68,7 +68,7 @@ void nvfp4_multi_tensor_2d_partial_cast(std::vector inp_list, return; } - auto stream = at::cuda::getCurrentCUDAStream(); + auto stream = getCurrentCUDAStream(); for (size_t i = 0; i < num_tensors; ++i) { const auto& inp = inp_list[i]; @@ -80,12 +80,12 @@ void nvfp4_multi_tensor_2d_partial_cast(std::vector inp_list, const size_t start_offset = static_cast(start_offset_list[i]); TORCH_CHECK(scale.dim() == 2, "scale must be a 2D tensor"); - TORCH_CHECK(scale.scalar_type() == at::ScalarType::Float, "scale must be a float tensor"); + TORCH_CHECK(scale.scalar_type() == ScalarType::Float, "scale must be a float tensor"); TORCH_CHECK(global_scale.numel() == 1, "global_scale must be a scalar tensor"); - TORCH_CHECK(global_scale.scalar_type() == at::ScalarType::Float, + TORCH_CHECK(global_scale.scalar_type() == ScalarType::Float, "global_scale must be a float tensor"); TORCH_CHECK( - inp.scalar_type() == at::ScalarType::Float || inp.scalar_type() == at::ScalarType::BFloat16, + inp.scalar_type() == ScalarType::Float || inp.scalar_type() == ScalarType::BFloat16, "input must be a float or bfloat16 tensor"); const TensorWrapper inp_cu = makeTransformerEngineTensor(inp.contiguous()); @@ -100,8 +100,8 @@ void nvfp4_multi_tensor_2d_partial_cast(std::vector inp_list, } void nvfp4_multi_tensor_compute_partial_amax( - std::vector master_weight_list, std::vector partial_amax_list, - std::vector global_amax_list, std::vector h_list, + std::vector master_weight_list, std::vector partial_amax_list, + std::vector global_amax_list, std::vector h_list, std::vector w_list, std::vector start_offset_list, int64_t block_len) { TORCH_CHECK(block_len == 16, "Currently only block_len = 16 is supported for NVFP4 2D"); @@ -116,7 +116,7 @@ void nvfp4_multi_tensor_compute_partial_amax( return; } - auto stream = at::cuda::getCurrentCUDAStream(); + auto stream = getCurrentCUDAStream(); for (size_t i = 0; i < num_tensors; ++i) { const auto& master_weight = master_weight_list[i]; @@ -127,12 +127,12 @@ void nvfp4_multi_tensor_compute_partial_amax( const size_t start_offset = static_cast(start_offset_list[i]); TORCH_CHECK(partial_amax.dim() == 2, "partial_amax must be a 2D tensor"); - TORCH_CHECK(partial_amax.scalar_type() == at::ScalarType::Float, + TORCH_CHECK(partial_amax.scalar_type() == ScalarType::Float, "partial_amax must be a float tensor"); - TORCH_CHECK(master_weight.scalar_type() == at::ScalarType::Float || - master_weight.scalar_type() == at::ScalarType::BFloat16, + TORCH_CHECK(master_weight.scalar_type() == ScalarType::Float || + master_weight.scalar_type() == ScalarType::BFloat16, "master_weight must be a float or bfloat16 tensor"); - TORCH_CHECK(global_amax.scalar_type() == at::ScalarType::Float, + TORCH_CHECK(global_amax.scalar_type() == ScalarType::Float, "global_amax must be a float tensor"); TORCH_CHECK(global_amax.numel() == 1, "global_amax must have exactly one element"); diff --git a/transformer_engine/pytorch/csrc/extensions/nvshmem_comm.cpp b/transformer_engine/pytorch/csrc/extensions/nvshmem_comm.cpp index ac68727ac8..84778ce8ee 100644 --- a/transformer_engine/pytorch/csrc/extensions/nvshmem_comm.cpp +++ b/transformer_engine/pytorch/csrc/extensions/nvshmem_comm.cpp @@ -14,12 +14,10 @@ #include #include -#include -#include namespace transformer_engine::pytorch { -void init_nvshmem_backend(c10d::ProcessGroup *process_group) { +void init_nvshmem_backend(ProcessGroup *process_group) { #ifdef NVTE_ENABLE_NVSHMEM nvshmemx_init_attr_t attr = {}; nvshmemx_uniqueid_t id = {}; @@ -30,23 +28,23 @@ void init_nvshmem_backend(c10d::ProcessGroup *process_group) { nvshmemx_get_uniqueid(&id); } - auto backend_is_nccl = (process_group->getBackendType() == c10d::ProcessGroup::BackendType::NCCL); + auto backend_is_nccl = (process_group->getBackendType() == ProcessGroup::BackendType::NCCL); NVTE_CHECK(backend_is_nccl, "Currently only support NCCL boostrap for NVSHMEM"); auto datatensor = - torch::from_blob(reinterpret_cast(&id), + from_blob(reinterpret_cast(&id), {static_cast(sizeof(nvshmemx_uniqueid_t) / sizeof(uint8_t))}, - at::device(torch::kCPU).dtype(torch::kUInt8)); + TensorOptions().device(kCPU).dtype(kUInt8)); auto datatmp = (backend_is_nccl) ? datatensor.cuda() : datatensor; - c10d::BroadcastOptions bcast_opts; + BroadcastOptions bcast_opts; bcast_opts.rootRank = 0; - std::vector datachunk = {datatmp}; + std::vector datachunk = {datatmp}; auto work = process_group->broadcast(datachunk, bcast_opts); work->wait(); if (backend_is_nccl) { datatensor.copy_(datatmp.cpu()); - datatmp = torch::Tensor(); + datatmp = Tensor(); } nvshmemx_set_attr_uniqueid_args(my_rank, num_ranks, &id, &attr); @@ -62,10 +60,10 @@ void init_nvshmem_backend(c10d::ProcessGroup *process_group) { #endif } -void nvshmem_wait_on_current_stream(torch::Tensor signal, const std::string &wait_kind) { +void nvshmem_wait_on_current_stream(Tensor signal, const std::string &wait_kind) { #ifdef NVTE_ENABLE_NVSHMEM uint64_t *sig_addr = reinterpret_cast(signal.data_ptr()); - cudaStream_t cur_stream = (cudaStream_t)at::cuda::getCurrentCUDAStream(); + cudaStream_t cur_stream = (cudaStream_t)getCurrentCUDAStream(); WaitKind wait_kind_enum = WaitKind::STREAM_WAIT; @@ -87,13 +85,13 @@ void nvshmem_wait_on_current_stream(torch::Tensor signal, const std::string &wai #endif } -torch::Tensor create_nvshmem_tensor(const std::vector &shape, c10::ScalarType dtype) { +Tensor create_nvshmem_tensor(const std::vector &shape, ScalarType dtype) { #ifdef NVTE_ENABLE_NVSHMEM auto option_gpu = - at::TensorOptions().dtype(dtype).device(at::kCUDA).device_index(c10::cuda::current_device()); - auto size = torch::elementSize(dtype) * + TensorOptions().dtype(dtype).device(kCUDA).device_index(current_device()); + auto size = elementSize(dtype) * std::accumulate(shape.begin(), shape.end(), 1, std::multiplies<>()); - return at::from_blob( + return from_blob( nvshmem_malloc(size), shape, [](void *ptr) { nvshmem_free(ptr); }, option_gpu); #else NVTE_ERROR("Internal TE error: create_nvshmem_tensor cannot be initialized with valid PyTorch ", @@ -101,15 +99,15 @@ torch::Tensor create_nvshmem_tensor(const std::vector &shape, c10::Scal #endif } -void nvshmem_send_on_current_stream(torch::Tensor src, torch::Tensor dst, int peer, - torch::Tensor signal) { +void nvshmem_send_on_current_stream(Tensor src, Tensor dst, int peer, + Tensor signal) { #ifdef NVTE_ENABLE_NVSHMEM void *src_ptr = reinterpret_cast(src.data_ptr()); void *dst_ptr = reinterpret_cast(dst.data_ptr()); uint64_t *sig_addr = reinterpret_cast(signal.data_ptr()); auto nelement = src.numel() * src.element_size(); uint64_t sigval = 1; - at::cuda::CUDAStream cur_stream = at::cuda::getCurrentCUDAStream(); + CUDAStream cur_stream = getCurrentCUDAStream(); nvshmemx_putmem_signal_on_stream(dst_ptr, src_ptr, nelement, sig_addr, sigval, NVSHMEM_SIGNAL_SET, peer, (cudaStream_t)cur_stream); diff --git a/transformer_engine/pytorch/csrc/extensions/padding.cpp b/transformer_engine/pytorch/csrc/extensions/padding.cpp index 6c66fda015..ba1ad5d522 100644 --- a/transformer_engine/pytorch/csrc/extensions/padding.cpp +++ b/transformer_engine/pytorch/csrc/extensions/padding.cpp @@ -9,7 +9,7 @@ namespace transformer_engine::pytorch { -void fused_multi_row_padding(at::Tensor input, at::Tensor output, +void fused_multi_row_padding(Tensor input, Tensor output, std::vector input_row_list, std::vector padded_input_row_list) { NVTE_CHECK(input_row_list.size() == padded_input_row_list.size(), @@ -77,11 +77,11 @@ void fused_multi_row_padding(at::Tensor input, at::Tensor output, // Launch TE kernel NVTE_SCOPED_GIL_RELEASE({ nvte_multi_padding(nvte_input_list.size(), nvte_input_list.data(), nvte_output_list.data(), - padded_num_rows_list.data(), at::cuda::getCurrentCUDAStream()); + padded_num_rows_list.data(), getCurrentCUDAStream()); }); } -void fused_multi_row_unpadding(at::Tensor input, at::Tensor output, +void fused_multi_row_unpadding(Tensor input, Tensor output, std::vector input_row_list, std::vector unpadded_input_row_list) { using namespace transformer_engine; @@ -151,7 +151,7 @@ void fused_multi_row_unpadding(at::Tensor input, at::Tensor output, // Launch TE kernel nvte_multi_unpadding(nvte_input_list.size(), nvte_input_list.data(), nvte_output_list.data(), - unpadded_num_rows_list.data(), at::cuda::getCurrentCUDAStream()); + unpadded_num_rows_list.data(), getCurrentCUDAStream()); } } // namespace transformer_engine::pytorch diff --git a/transformer_engine/pytorch/csrc/extensions/permutation.cpp b/transformer_engine/pytorch/csrc/extensions/permutation.cpp index 226705b169..435bf0202d 100644 --- a/transformer_engine/pytorch/csrc/extensions/permutation.cpp +++ b/transformer_engine/pytorch/csrc/extensions/permutation.cpp @@ -8,9 +8,9 @@ namespace transformer_engine::pytorch { -std::tuple> moe_permute_fwd( - at::Tensor input, const DType dtype, at::Tensor indices, int64_t num_out_tokens, - std::vector workspace, int64_t max_expanded_token_num) { +std::tuple> moe_permute_fwd( + Tensor input, const DType dtype, Tensor indices, int64_t num_out_tokens, + std::vector workspace, int64_t max_expanded_token_num) { const int num_tokens = input.size(0); int num_cols = input.size(1); const int topK = indices.size(1); @@ -18,19 +18,19 @@ std::tuple> moe_permute_fwd( // Initialize the workspace on the first run if (workspace.empty()) { auto options = - torch::TensorOptions().dtype(torch::kInt32).device(torch::kCUDA).requires_grad(false); + TensorOptions().dtype(kInt32).device(kCUDA).requires_grad(false); - at::Tensor sorted_indices = torch::empty(max_expanded_token_num, options); - at::Tensor row_id = torch::range(0, max_expanded_token_num - 1, 1, options); - at::Tensor sorted_row_id = - torch::empty(max_expanded_token_num, - torch::dtype(torch::kInt32).device(torch::kCUDA).requires_grad(false)); + Tensor sorted_indices = empty(max_expanded_token_num, options); + Tensor row_id = range(0, max_expanded_token_num - 1, 1, options); + Tensor sorted_row_id = + empty(max_expanded_token_num, + TensorOptions().dtype(kInt32).device(kCUDA).requires_grad(false)); size_t temp_storage_bytes = 0; nvte_device_radix_sort_pairs(nullptr, &temp_storage_bytes, nullptr, nullptr, nullptr, nullptr, max_expanded_token_num); - at::Tensor temp_storage = torch::empty( - temp_storage_bytes, torch::dtype(torch::kInt8).device(torch::kCUDA).requires_grad(false)); + Tensor temp_storage = empty( + temp_storage_bytes, TensorOptions().dtype(kInt8).device(kCUDA).requires_grad(false)); workspace.push_back(sorted_indices); workspace.push_back(row_id); @@ -53,13 +53,13 @@ std::tuple> moe_permute_fwd( // Output buffer alloc num_out_tokens = (num_out_tokens > 0) ? num_out_tokens : num_tokens * topK; - at::Tensor permuted_output = - torch::empty({num_out_tokens, num_cols}, - torch::dtype(input.scalar_type()).device(torch::kCUDA).requires_grad(false)); - at::Tensor row_id_map = torch::empty( - {num_tokens * topK}, torch::dtype(torch::kInt32).device(torch::kCUDA).requires_grad(false)); + Tensor permuted_output = + empty({num_out_tokens, num_cols}, + TensorOptions().dtype(input.scalar_type()).device(kCUDA).requires_grad(false)); + Tensor row_id_map = empty( + {num_tokens * topK}, TensorOptions().dtype(kInt32).device(kCUDA).requires_grad(false)); - auto stream = at::cuda::getCurrentCUDAStream().stream(); + auto stream = getCurrentCUDAStream().stream(); auto input_cu = makeTransformerEngineTensor( input.data_ptr(), @@ -82,21 +82,21 @@ std::tuple> moe_permute_fwd( return std::make_tuple(permuted_output, row_id_map, workspace); } -at::Tensor moe_permute_bwd(at::Tensor input, const DType dtype, at::Tensor row_id_map, - at::Tensor prob, int64_t num_tokens, int64_t topK) { +Tensor moe_permute_bwd(Tensor input, const DType dtype, Tensor row_id_map, + Tensor prob, int64_t num_tokens, int64_t topK) { return moe_unpermute_fwd(input, dtype, row_id_map, prob, num_tokens, topK); } -at::Tensor moe_unpermute_fwd(at::Tensor input, const DType dtype, at::Tensor row_id_map, - at::Tensor prob, int64_t num_tokens, int64_t topK) { +Tensor moe_unpermute_fwd(Tensor input, const DType dtype, Tensor row_id_map, + Tensor prob, int64_t num_tokens, int64_t topK) { int num_cols = input.size(1); // Output buffer alloc - at::Tensor unpermuted_output = - torch::empty({num_tokens, num_cols}, - torch::dtype(input.scalar_type()).device(torch::kCUDA).requires_grad(false)); + Tensor unpermuted_output = + empty({num_tokens, num_cols}, + TensorOptions().dtype(input.scalar_type()).device(kCUDA).requires_grad(false)); - auto stream = at::cuda::getCurrentCUDAStream().stream(); + auto stream = getCurrentCUDAStream().stream(); auto input_cu = makeTransformerEngineTensor( input.data_ptr(), @@ -116,21 +116,21 @@ at::Tensor moe_unpermute_fwd(at::Tensor input, const DType dtype, at::Tensor row return unpermuted_output; } -std::tuple moe_unpermute_bwd(at::Tensor input_bwd, at::Tensor input_fwd, - const DType dtype, at::Tensor row_id_map, - at::Tensor prob) { +std::tuple moe_unpermute_bwd(Tensor input_bwd, Tensor input_fwd, + const DType dtype, Tensor row_id_map, + Tensor prob) { const int topK = (prob.numel() > 0) ? prob.size(1) : 1; const int num_tokens = (prob.numel() > 0) ? prob.size(0) : row_id_map.size(0); int num_cols = input_bwd.size(1); // Output buffer alloc - at::Tensor act_grad = - torch::empty({input_fwd.size(0), num_cols}, - torch::dtype(input_bwd.scalar_type()).device(torch::kCUDA).requires_grad(false)); - at::Tensor prob_grad = torch::empty( - {num_tokens, topK}, torch::dtype(torch::kFloat32).device(torch::kCUDA).requires_grad(false)); + Tensor act_grad = + empty({input_fwd.size(0), num_cols}, + TensorOptions().dtype(input_bwd.scalar_type()).device(kCUDA).requires_grad(false)); + Tensor prob_grad = empty( + {num_tokens, topK}, TensorOptions().dtype(kFloat32).device(kCUDA).requires_grad(false)); - auto stream = at::cuda::getCurrentCUDAStream().stream(); + auto stream = getCurrentCUDAStream().stream(); auto input_bwd_cu = makeTransformerEngineTensor( input_bwd.data_ptr(), diff --git a/transformer_engine/pytorch/csrc/extensions/pybind.cpp b/transformer_engine/pytorch/csrc/extensions/pybind.cpp index d6089b1e01..a9f455de44 100644 --- a/transformer_engine/pytorch/csrc/extensions/pybind.cpp +++ b/transformer_engine/pytorch/csrc/extensions/pybind.cpp @@ -177,6 +177,10 @@ void bind_quantize_with_amax_extensions(py::module_ &m) { #include "common/util/pybind_helper.h" +// PYBIND11_MODULE below is at global scope; bring the facade aliases +// (Tensor, ProcessGroup, ...) into scope for the binding definitions. +using namespace transformer_engine::pytorch; // NOLINT(build/namespaces) + PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { NVTE_DECLARE_COMMON_PYBIND11_HANDLES(m) @@ -488,7 +492,7 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { m.def("multi_tensor_transpose_to_bhsd", &transformer_engine::pytorch::multi_tensor_transpose_to_bhsd, "Permute multiple tensors from BSHD/SBHD to BHSD.", py::arg("inputs"), - py::arg("original_format"), py::arg("outputs") = std::vector>{}, + py::arg("original_format"), py::arg("outputs") = std::vector>{}, py::call_guard()); m.def("multi_tensor_pad_last_dim", &transformer_engine::pytorch::multi_tensor_pad_last_dim, "Pad multiple tensors' last dimension to a common alignment.", py::arg("inputs"), @@ -683,13 +687,13 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { py::class_(m, "CommOverlapHelper") .def(py::init<>(), py::call_guard()) - .def(py::init>(), + .def(py::init>(), py::call_guard(), py::arg("world_group"), py::arg("intra_node_group") = py::none()); py::class_, transformer_engine::CommOverlapBase, transformer_engine::CommOverlapCore>(m, "CommOverlap") - .def(py::init([](const std::vector &buffer_shape, at::ScalarType buffer_dtype, + .def(py::init([](const std::vector &buffer_shape, ScalarType buffer_dtype, CommOverlapHelper *helper, int tp_size, bool use_cublasmp, transformer_engine::CommOverlapType comm_type, int num_splits, int num_max_streams, int comm_cga_size, int gemm_priority, int comm_priority, @@ -714,7 +718,7 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { py::arg("num_comm_sm") = 16, py::arg("set_sm_margin") = true, py::arg("atomic_gemm") = false, py::arg("rs_overlap_first_gemm") = false) .def("copy_into_buffer", - static_cast( + static_cast( &CommOverlap::copy_into_buffer), py::arg("input"), py::arg("local_chunk") = false) .def("get_buffer", &CommOverlap::get_buffer, py::arg("local_chunk") = false, @@ -724,7 +728,7 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { py::class_, transformer_engine::CommOverlapP2PBase, transformer_engine::CommOverlapCore>( m, "CommOverlapP2P") - .def(py::init([](const std::vector &buffer_shape, at::ScalarType buffer_dtype, + .def(py::init([](const std::vector &buffer_shape, ScalarType buffer_dtype, CommOverlapHelper *helper, int tp_size, transformer_engine::CommOverlapType comm_type, int num_max_streams, int comm_cga_size, int gemm_priority, int comm_priority, int num_comm_sm, @@ -747,7 +751,7 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { py::arg("set_sm_margin") = false, py::arg("atomic_gemm") = false, py::arg("use_ce") = true, py::arg("aggregate") = false, py::arg("use_cublasmp") = false) .def("copy_into_buffer", - static_cast( + static_cast( &CommOverlapP2P::copy_into_buffer), py::arg("input"), py::arg("local_chunk") = false) .def("get_buffer", &CommOverlapP2P::get_buffer, py::arg("local_chunk") = false, diff --git a/transformer_engine/pytorch/csrc/extensions/recipe.cpp b/transformer_engine/pytorch/csrc/extensions/recipe.cpp index 9be288d2f7..c19e357487 100644 --- a/transformer_engine/pytorch/csrc/extensions/recipe.cpp +++ b/transformer_engine/pytorch/csrc/extensions/recipe.cpp @@ -4,9 +4,6 @@ * See LICENSE for license information. ************************************************************************/ -#include -#include - #include #include "../extensions.h" @@ -14,11 +11,11 @@ namespace transformer_engine::pytorch { -void compute_amax(const at::Tensor& tensor, at::Tensor& amax) { +void compute_amax(const Tensor& tensor, Tensor& amax) { auto input_tensor = tensor.contiguous(); const TensorWrapper& te_input = makeTransformerEngineTensor(input_tensor); - TORCH_CHECK(amax.scalar_type() == at::kFloat, "amax must be a float tensor"); + TORCH_CHECK(amax.scalar_type() == kFloat, "amax must be a float tensor"); TORCH_CHECK(amax.numel() == 1, "amax must have exactly one element"); auto* amax_ptr = amax.data_ptr(); TensorWrapper fake_te_output( @@ -26,12 +23,12 @@ void compute_amax(const at::Tensor& tensor, at::Tensor& amax) { DType::kFloat32, // It doesn't matter because we only compute amax. amax_ptr); - nvte_compute_amax(te_input.data(), fake_te_output.data(), at::cuda::getCurrentCUDAStream()); + nvte_compute_amax(te_input.data(), fake_te_output.data(), getCurrentCUDAStream()); } -void fused_amax_and_scale_update_after_reduction(const at::Tensor& amax_reduction_buffer, - std::vector amax_histories, - std::vector scales, +void fused_amax_and_scale_update_after_reduction(const Tensor& amax_reduction_buffer, + std::vector amax_histories, + std::vector scales, const std::string& amax_compute_algo, DType fp8_dtype, float margin) { size_t num_tensors = amax_histories.size(); @@ -58,7 +55,7 @@ void fused_amax_and_scale_update_after_reduction(const at::Tensor& amax_reductio makeTransformerEngineTensor(amax_reduction_buffer).data(), std::vector(te_amax_histories.begin(), te_amax_histories.end()), std::vector(te_scales.begin(), te_scales.end()), amax_compute_algo.c_str(), - static_cast(fp8_dtype), margin, at::cuda::getCurrentCUDAStream()); + static_cast(fp8_dtype), margin, getCurrentCUDAStream()); } } // namespace transformer_engine::pytorch diff --git a/transformer_engine/pytorch/csrc/extensions/router.cpp b/transformer_engine/pytorch/csrc/extensions/router.cpp index 5dc5c7fe86..71810b4fb6 100644 --- a/transformer_engine/pytorch/csrc/extensions/router.cpp +++ b/transformer_engine/pytorch/csrc/extensions/router.cpp @@ -17,21 +17,21 @@ static std::map score_function_map = { // Allocate a routing_map output tensor: // BYTEMAP -> bool [*leading_dims, num_experts] // BITMAP_U8 -> uint8[*leading_dims, ceil(num_experts/8)], LSB-first -static at::Tensor allocate_routing_map(c10::IntArrayRef leading_dims, int64_t num_experts, +static Tensor allocate_routing_map(IntArrayRef leading_dims, int64_t num_experts, int routing_map_format) { std::vector shape(leading_dims.begin(), leading_dims.end()); if (routing_map_format == NVTE_ROUTING_MAP_FORMAT_BITMAP_U8) { shape.push_back((num_experts + 7) / 8); - return at::empty(shape, at::dtype(at::kByte).device(at::kCUDA)); + return empty(shape, TensorOptions().dtype(kByte).device(kCUDA)); } shape.push_back(num_experts); - return at::empty(shape, at::dtype(at::kBool).device(at::kCUDA)); + return empty(shape, TensorOptions().dtype(kBool).device(kCUDA)); } -std::tuple fused_topk_with_score_function_fwd( - at::Tensor logits, int topk, bool use_pre_softmax, std::optional num_groups, +std::tuple fused_topk_with_score_function_fwd( + Tensor logits, int topk, bool use_pre_softmax, std::optional num_groups, std::optional group_topk, std::optional scaling_factor, std::string score_function, - std::optional expert_bias, int routing_map_format) { + std::optional expert_bias, int routing_map_format) { TORCH_CHECK(logits.dim() >= 1, "logits must have at least 1 dim"); TORCH_CHECK(logits.is_contiguous(), "logits must be contiguous"); auto sizes = logits.sizes(); @@ -44,7 +44,7 @@ std::tuple fused_topk_with_score_function_fw if (expert_bias.has_value()) { TORCH_CHECK(score_function == "sigmoid" || score_function == "sqrtsoftplus", "score_function must be sigmoid or sqrtsoftplus when expert_bias is not None"); - TORCH_CHECK(expert_bias.value().scalar_type() == at::kFloat, + TORCH_CHECK(expert_bias.value().scalar_type() == kFloat, "expert_bias must be a float32 tensor"); } // Check if the score function is valid @@ -60,10 +60,10 @@ std::tuple fused_topk_with_score_function_fw int num_groups_value = num_groups.has_value() ? num_groups.value() : -1; float scaling_factor_value = scaling_factor.has_value() ? scaling_factor.value() : 1.0f; - at::Tensor probs = at::empty(sizes, at::dtype(logits.scalar_type()).device(at::kCUDA)); - at::Tensor routing_map = + Tensor probs = empty(sizes, TensorOptions().dtype(logits.scalar_type()).device(kCUDA)); + Tensor routing_map = allocate_routing_map(sizes.slice(0, sizes.size() - 1), num_experts, routing_map_format); - at::Tensor intermediate_output = at::empty(sizes, at::dtype(at::kFloat).device(at::kCUDA)); + Tensor intermediate_output = empty(sizes, TensorOptions().dtype(kFloat).device(kCUDA)); // 2D shape for the kernel (common-layer NVTE_CHECKs require {num_tokens, trailing_dim}). const std::vector shape_2d = {static_cast(num_tokens), @@ -92,13 +92,13 @@ std::tuple fused_topk_with_score_function_fw use_pre_softmax, num_groups_value, group_topk_value, scaling_factor_value, score_function_map[score_function], expert_bias_cu.data(), probs_cu.data(), routing_map_cu.data(), static_cast(routing_map_format), - intermediate_output_cu.data(), at::cuda::getCurrentCUDAStream()); + intermediate_output_cu.data(), getCurrentCUDAStream()); return std::make_tuple(probs, routing_map, intermediate_output); } -void fused_topk_with_score_function_bwd(at::Tensor routing_map, at::Tensor intermediate_output, - at::Tensor grad_probs, at::Tensor grad_logits, int topk, +void fused_topk_with_score_function_bwd(Tensor routing_map, Tensor intermediate_output, + Tensor grad_probs, Tensor grad_logits, int topk, bool use_pre_softmax, std::optional scaling_factor, std::string score_function, int routing_map_format) { TORCH_CHECK(grad_probs.dim() >= 1, "grad_probs must have at least 1 dim"); @@ -133,11 +133,11 @@ void fused_topk_with_score_function_bwd(at::Tensor routing_map, at::Tensor inter routing_map_cu.data(), static_cast(routing_map_format), intermediate_output_cu.data(), grad_probs_cu.data(), static_cast(num_tokens), static_cast(num_experts), topk, use_pre_softmax, scaling_factor_value, - score_function_value, grad_logits_cu.data(), at::cuda::getCurrentCUDAStream()); + score_function_value, grad_logits_cu.data(), getCurrentCUDAStream()); } -std::tuple fused_score_for_moe_aux_loss_fwd( - at::Tensor logits, int topk, std::string score_function, int routing_map_format) { +std::tuple fused_score_for_moe_aux_loss_fwd( + Tensor logits, int topk, std::string score_function, int routing_map_format) { TORCH_CHECK(logits.dim() >= 1, "logits must have at least 1 dim"); TORCH_CHECK(logits.is_contiguous(), "logits must be contiguous"); auto sizes = logits.sizes(); @@ -152,10 +152,10 @@ std::tuple fused_score_for_moe_aux_loss_fwd( "score_function must be softmax, sigmoid or sqrtsoftplus for router fusion"); int score_function_value = score_function_map[score_function]; - at::Tensor scores = at::empty(sizes, at::dtype(at::kFloat).device(at::kCUDA)); - at::Tensor routing_map = + Tensor scores = empty(sizes, TensorOptions().dtype(kFloat).device(kCUDA)); + Tensor routing_map = allocate_routing_map(sizes.slice(0, sizes.size() - 1), num_experts, routing_map_format); - at::Tensor intermediate_output = at::empty(sizes, at::dtype(at::kFloat).device(at::kCUDA)); + Tensor intermediate_output = empty(sizes, TensorOptions().dtype(kFloat).device(kCUDA)); const std::vector shape_2d = {static_cast(num_tokens), static_cast(num_experts)}; @@ -178,13 +178,13 @@ std::tuple fused_score_for_moe_aux_loss_fwd( logits_cu.data(), static_cast(num_tokens), static_cast(num_experts), topk, score_function_value, scores_cu.data(), routing_map_cu.data(), static_cast(routing_map_format), intermediate_output_cu.data(), - at::cuda::getCurrentCUDAStream()); + getCurrentCUDAStream()); return std::make_tuple(scores, routing_map, intermediate_output); } -void fused_score_for_moe_aux_loss_bwd(at::Tensor intermediate_output, at::Tensor grad_scores, - at::Tensor grad_logits, int topk, +void fused_score_for_moe_aux_loss_bwd(Tensor intermediate_output, Tensor grad_scores, + Tensor grad_logits, int topk, std::string score_function) { TORCH_CHECK(grad_scores.dim() >= 1, "grad_scores must have at least 1 dim"); TORCH_CHECK(grad_scores.is_contiguous(), "grad_scores must be contiguous"); @@ -210,11 +210,11 @@ void fused_score_for_moe_aux_loss_bwd(at::Tensor intermediate_output, at::Tensor nvte_fused_score_for_moe_aux_loss_backward( intermediate_output_cu.data(), grad_scores_cu.data(), static_cast(num_tokens), static_cast(num_experts), topk, score_function_value, grad_logits_cu.data(), - at::cuda::getCurrentCUDAStream()); + getCurrentCUDAStream()); } -std::tuple fused_moe_aux_loss_fwd(at::Tensor probs, - at::Tensor tokens_per_expert, +std::tuple fused_moe_aux_loss_fwd(Tensor probs, + Tensor tokens_per_expert, int total_num_tokens, int num_experts, int num_rows, int num_cols, int topk, float coeff) { @@ -223,8 +223,8 @@ std::tuple fused_moe_aux_loss_fwd(at::Tensor probs, TORCH_CHECK(num_experts > 0, "num_experts must be greater than 0"); // Create the output tensor - at::Tensor aux_loss = at::empty({}, at::dtype(probs.scalar_type()).device(at::kCUDA)); - at::Tensor Const_buf = at::empty({2}, at::dtype(at::kFloat).device(at::kCUDA)); + Tensor aux_loss = empty({}, TensorOptions().dtype(probs.scalar_type()).device(kCUDA)); + Tensor Const_buf = empty({2}, TensorOptions().dtype(kFloat).device(kCUDA)); auto probs_cu = makeTransformerEngineTensor(probs); auto tokens_per_expert_cu = makeTransformerEngineTensor(tokens_per_expert); @@ -233,16 +233,16 @@ std::tuple fused_moe_aux_loss_fwd(at::Tensor probs, nvte_fused_moe_aux_loss_forward(probs_cu.data(), tokens_per_expert_cu.data(), total_num_tokens, num_experts, num_rows, num_cols, topk, coeff, aux_loss_cu.data(), - Const_buf_cu.data(), at::cuda::getCurrentCUDAStream()); + Const_buf_cu.data(), getCurrentCUDAStream()); return std::make_tuple(aux_loss, Const_buf); } -at::Tensor fused_moe_aux_loss_bwd(at::Tensor Const_buf, at::Tensor tokens_per_expert, int num_rows, - int num_cols, at::Tensor grad_aux_loss) { +Tensor fused_moe_aux_loss_bwd(Tensor Const_buf, Tensor tokens_per_expert, int num_rows, + int num_cols, Tensor grad_aux_loss) { // Create the output tensor - at::Tensor grad_probs = - at::empty({num_rows, num_cols}, at::dtype(grad_aux_loss.scalar_type()).device(at::kCUDA)); + Tensor grad_probs = + empty({num_rows, num_cols}, TensorOptions().dtype(grad_aux_loss.scalar_type()).device(kCUDA)); auto Const_buf_cu = makeTransformerEngineTensor(Const_buf); auto tokens_per_expert_cu = makeTransformerEngineTensor(tokens_per_expert); @@ -252,7 +252,7 @@ at::Tensor fused_moe_aux_loss_bwd(at::Tensor Const_buf, at::Tensor tokens_per_ex // Meta data for the kernel nvte_fused_moe_aux_loss_backward(Const_buf_cu.data(), tokens_per_expert_cu.data(), num_rows, num_cols, grad_aux_loss_cu.data(), grad_probs_cu.data(), - at::cuda::getCurrentCUDAStream()); + getCurrentCUDAStream()); return grad_probs; } diff --git a/transformer_engine/pytorch/csrc/extensions/softmax.cpp b/transformer_engine/pytorch/csrc/extensions/softmax.cpp index 3bb6a5e7b3..ecea67f9a2 100644 --- a/transformer_engine/pytorch/csrc/extensions/softmax.cpp +++ b/transformer_engine/pytorch/csrc/extensions/softmax.cpp @@ -8,10 +8,10 @@ namespace transformer_engine::pytorch { -at::Tensor scaled_softmax_forward(at::Tensor input, float scale_factor) { +Tensor scaled_softmax_forward(Tensor input, float scale_factor) { AT_ASSERTM(input.dim() == 4, "expected 4D tensor"); - AT_ASSERTM((input.scalar_type() == at::ScalarType::Half) || - (input.scalar_type() == at::ScalarType::BFloat16), + AT_ASSERTM((input.scalar_type() == ScalarType::Half) || + (input.scalar_type() == ScalarType::BFloat16), "Only fp16 and bf16 are supported"); const int batches = input.size(0); @@ -26,18 +26,18 @@ at::Tensor scaled_softmax_forward(at::Tensor input, float scale_factor) { // Output auto act_options = input.options().requires_grad(false); auto softmax_results = - torch::empty({batches, attn_heads, query_seq_len, key_seq_len}, act_options); + empty({batches, attn_heads, query_seq_len, key_seq_len}, act_options); auto input_cu = makeTransformerEngineTensor(input); auto softmax_results_cu = makeTransformerEngineTensor(softmax_results); nvte_scaled_softmax_forward(input_cu.data(), softmax_results_cu.data(), scale_factor, - at::cuda::getCurrentCUDAStream()); + getCurrentCUDAStream()); return softmax_results; } -at::Tensor scaled_softmax_backward(at::Tensor output_grad_, at::Tensor softmax_results_, +Tensor scaled_softmax_backward(Tensor output_grad_, Tensor softmax_results_, float scale_factor) { auto output_grads = output_grad_.contiguous(); auto softmax_results = softmax_results_.contiguous(); @@ -45,11 +45,11 @@ at::Tensor scaled_softmax_backward(at::Tensor output_grad_, at::Tensor softmax_r AT_ASSERTM(output_grads.dim() == 4, "expected 4D tensor"); AT_ASSERTM(softmax_results.dim() == 4, "expected 4D tensor"); - AT_ASSERTM((output_grads.scalar_type() == at::ScalarType::Half) || - (output_grads.scalar_type() == at::ScalarType::BFloat16), + AT_ASSERTM((output_grads.scalar_type() == ScalarType::Half) || + (output_grads.scalar_type() == ScalarType::BFloat16), "Only fp16 and bf16 are supported"); - AT_ASSERTM((softmax_results.scalar_type() == at::ScalarType::Half) || - (softmax_results.scalar_type() == at::ScalarType::BFloat16), + AT_ASSERTM((softmax_results.scalar_type() == ScalarType::Half) || + (softmax_results.scalar_type() == ScalarType::BFloat16), "Only fp16 and bf16 are supported"); auto output_grads_cu = makeTransformerEngineTensor(output_grads); @@ -58,15 +58,15 @@ at::Tensor scaled_softmax_backward(at::Tensor output_grad_, at::Tensor softmax_r // Produce gradients in place. nvte_scaled_softmax_backward(output_grads_cu.data(), softmax_results_cu.data(), output_grads_cu.data(), scale_factor, - at::cuda::getCurrentCUDAStream()); + getCurrentCUDAStream()); return output_grads; } -at::Tensor scaled_masked_softmax_forward(at::Tensor input, at::Tensor mask, float scale_factor) { +Tensor scaled_masked_softmax_forward(Tensor input, Tensor mask, float scale_factor) { AT_ASSERTM(input.dim() == 4, "expected 4D tensor"); - AT_ASSERTM((input.scalar_type() == at::ScalarType::Half) || - (input.scalar_type() == at::ScalarType::BFloat16), + AT_ASSERTM((input.scalar_type() == ScalarType::Half) || + (input.scalar_type() == ScalarType::BFloat16), "Only fp16 and bf16 are supported"); AT_ASSERTM(mask.dim() == 4, "expected 4D tensor"); if (!input.is_contiguous()) input = input.contiguous(); @@ -88,19 +88,19 @@ at::Tensor scaled_masked_softmax_forward(at::Tensor input, at::Tensor mask, floa auto act_options = input.options().requires_grad(false); auto softmax_results = - torch::empty({batches, attn_heads, query_seq_len, key_seq_len}, act_options); + empty({batches, attn_heads, query_seq_len, key_seq_len}, act_options); auto input_cu = makeTransformerEngineTensor(input); auto mask_cu = makeTransformerEngineTensor(mask); auto softmax_results_cu = makeTransformerEngineTensor(softmax_results); nvte_scaled_masked_softmax_forward(input_cu.data(), mask_cu.data(), softmax_results_cu.data(), - scale_factor, at::cuda::getCurrentCUDAStream()); + scale_factor, getCurrentCUDAStream()); return softmax_results; } -at::Tensor scaled_masked_softmax_backward(at::Tensor output_grad_, at::Tensor softmax_results_, +Tensor scaled_masked_softmax_backward(Tensor output_grad_, Tensor softmax_results_, float scale_factor) { auto output_grads = output_grad_.contiguous(); auto softmax_results = softmax_results_.contiguous(); @@ -108,11 +108,11 @@ at::Tensor scaled_masked_softmax_backward(at::Tensor output_grad_, at::Tensor so AT_ASSERTM(output_grads.dim() == 4, "expected 3D tensor"); AT_ASSERTM(softmax_results.dim() == 4, "expected 3D tensor"); - AT_ASSERTM((output_grads.scalar_type() == at::ScalarType::Half) || - (output_grads.scalar_type() == at::ScalarType::BFloat16), + AT_ASSERTM((output_grads.scalar_type() == ScalarType::Half) || + (output_grads.scalar_type() == ScalarType::BFloat16), "Only fp16 and bf16 are supported"); - AT_ASSERTM((softmax_results.scalar_type() == at::ScalarType::Half) || - (softmax_results.scalar_type() == at::ScalarType::BFloat16), + AT_ASSERTM((softmax_results.scalar_type() == ScalarType::Half) || + (softmax_results.scalar_type() == ScalarType::BFloat16), "Only fp16 and bf16 are supported"); auto output_grads_cu = makeTransformerEngineTensor(output_grads); @@ -121,15 +121,15 @@ at::Tensor scaled_masked_softmax_backward(at::Tensor output_grad_, at::Tensor so // Produce gradients in place. nvte_scaled_softmax_backward(output_grads_cu.data(), softmax_results_cu.data(), output_grads_cu.data(), scale_factor, - at::cuda::getCurrentCUDAStream()); + getCurrentCUDAStream()); return output_grads; } -at::Tensor scaled_upper_triang_masked_softmax_forward(at::Tensor input, float scale_factor) { +Tensor scaled_upper_triang_masked_softmax_forward(Tensor input, float scale_factor) { AT_ASSERTM(input.dim() == 3, "expected 3D tensor"); - AT_ASSERTM((input.scalar_type() == at::ScalarType::Half) || - (input.scalar_type() == at::ScalarType::BFloat16), + AT_ASSERTM((input.scalar_type() == ScalarType::Half) || + (input.scalar_type() == ScalarType::BFloat16), "Only fp16 and bf16 are supported"); const int attn_batches = input.size(0); @@ -138,19 +138,19 @@ at::Tensor scaled_upper_triang_masked_softmax_forward(at::Tensor input, float sc // Output auto act_options = input.options().requires_grad(false); - auto softmax_results = torch::empty({attn_batches, seq_len, seq_len}, act_options); + auto softmax_results = empty({attn_batches, seq_len, seq_len}, act_options); auto input_cu = makeTransformerEngineTensor(input); auto softmax_results_cu = makeTransformerEngineTensor(softmax_results); nvte_scaled_upper_triang_masked_softmax_forward(input_cu.data(), softmax_results_cu.data(), - scale_factor, at::cuda::getCurrentCUDAStream()); + scale_factor, getCurrentCUDAStream()); return softmax_results; } -at::Tensor scaled_upper_triang_masked_softmax_backward(at::Tensor output_grads_, - at::Tensor softmax_results_, +Tensor scaled_upper_triang_masked_softmax_backward(Tensor output_grads_, + Tensor softmax_results_, float scale_factor) { auto output_grads = output_grads_.contiguous(); auto softmax_results = softmax_results_.contiguous(); @@ -158,11 +158,11 @@ at::Tensor scaled_upper_triang_masked_softmax_backward(at::Tensor output_grads_, AT_ASSERTM(output_grads.dim() == 3, "expected 3D tensor"); AT_ASSERTM(softmax_results.dim() == 3, "expected 3D tensor"); - AT_ASSERTM((output_grads.scalar_type() == at::ScalarType::Half) || - (output_grads.scalar_type() == at::ScalarType::BFloat16), + AT_ASSERTM((output_grads.scalar_type() == ScalarType::Half) || + (output_grads.scalar_type() == ScalarType::BFloat16), "Only fp16 and bf16 are supported"); - AT_ASSERTM((softmax_results.scalar_type() == at::ScalarType::Half) || - (softmax_results.scalar_type() == at::ScalarType::BFloat16), + AT_ASSERTM((softmax_results.scalar_type() == ScalarType::Half) || + (softmax_results.scalar_type() == ScalarType::BFloat16), "Only fp16 and bf16 are supported"); TORCH_CHECK(output_grads.size(1) == output_grads.size(2)); @@ -173,15 +173,15 @@ at::Tensor scaled_upper_triang_masked_softmax_backward(at::Tensor output_grads_, // Produce gradients in place. nvte_scaled_upper_triang_masked_softmax_backward( output_grads_cu.data(), softmax_results_cu.data(), output_grads_cu.data(), scale_factor, - at::cuda::getCurrentCUDAStream()); + getCurrentCUDAStream()); return output_grads; } -at::Tensor scaled_aligned_causal_masked_softmax_forward(at::Tensor input, float scale_factor) { +Tensor scaled_aligned_causal_masked_softmax_forward(Tensor input, float scale_factor) { AT_ASSERTM(input.dim() == 4, "expected 4D tensor"); - AT_ASSERTM((input.scalar_type() == at::ScalarType::Half) || - (input.scalar_type() == at::ScalarType::BFloat16), + AT_ASSERTM((input.scalar_type() == ScalarType::Half) || + (input.scalar_type() == ScalarType::BFloat16), "Only fp16 and bf16 are supported"); const int batches = input.size(0); @@ -196,19 +196,19 @@ at::Tensor scaled_aligned_causal_masked_softmax_forward(at::Tensor input, float // Output auto act_options = input.options().requires_grad(false); auto softmax_results = - torch::empty({batches, attn_heads, query_seq_len, key_seq_len}, act_options); + empty({batches, attn_heads, query_seq_len, key_seq_len}, act_options); auto input_cu = makeTransformerEngineTensor(input); auto softmax_results_cu = makeTransformerEngineTensor(softmax_results); nvte_scaled_aligned_causal_masked_softmax_forward(input_cu.data(), softmax_results_cu.data(), - scale_factor, at::cuda::getCurrentCUDAStream()); + scale_factor, getCurrentCUDAStream()); return softmax_results; } -at::Tensor scaled_aligned_causal_masked_softmax_backward(at::Tensor output_grad_, - at::Tensor softmax_results_, +Tensor scaled_aligned_causal_masked_softmax_backward(Tensor output_grad_, + Tensor softmax_results_, float scale_factor) { auto output_grads = output_grad_.contiguous(); auto softmax_results = softmax_results_.contiguous(); @@ -216,11 +216,11 @@ at::Tensor scaled_aligned_causal_masked_softmax_backward(at::Tensor output_grad_ AT_ASSERTM(output_grads.dim() == 4, "expected 4D tensor"); AT_ASSERTM(softmax_results.dim() == 4, "expected 4D tensor"); - AT_ASSERTM((output_grads.scalar_type() == at::ScalarType::Half) || - (output_grads.scalar_type() == at::ScalarType::BFloat16), + AT_ASSERTM((output_grads.scalar_type() == ScalarType::Half) || + (output_grads.scalar_type() == ScalarType::BFloat16), "Only fp16 and bf16 are supported"); - AT_ASSERTM((softmax_results.scalar_type() == at::ScalarType::Half) || - (softmax_results.scalar_type() == at::ScalarType::BFloat16), + AT_ASSERTM((softmax_results.scalar_type() == ScalarType::Half) || + (softmax_results.scalar_type() == ScalarType::BFloat16), "Only fp16 and bf16 are supported"); auto output_grads_cu = makeTransformerEngineTensor(output_grads); @@ -229,7 +229,7 @@ at::Tensor scaled_aligned_causal_masked_softmax_backward(at::Tensor output_grad_ // Produce gradients in place. nvte_scaled_aligned_causal_masked_softmax_backward( output_grads_cu.data(), softmax_results_cu.data(), output_grads_cu.data(), scale_factor, - at::cuda::getCurrentCUDAStream()); + getCurrentCUDAStream()); return output_grads; } diff --git a/transformer_engine/pytorch/csrc/extensions/swizzle.cpp b/transformer_engine/pytorch/csrc/extensions/swizzle.cpp index c90a7d6d0d..a613a4870c 100644 --- a/transformer_engine/pytorch/csrc/extensions/swizzle.cpp +++ b/transformer_engine/pytorch/csrc/extensions/swizzle.cpp @@ -44,7 +44,7 @@ bool is_empty_grouped_tensor_param(const NVTEBasicTensor &t) { } // namespace -std::tuple, std::optional> swizzle_scales_for_gemm( +std::tuple, std::optional> swizzle_scales_for_gemm( transformer_engine::TensorWrapper &tensor, bool rowwise_usage, bool columnwise_usage) { // Return early if scale swizzling is not required const auto scaling_mode = tensor.scaling_mode(); @@ -66,10 +66,10 @@ std::tuple, std::optional> swizzle_scales_ } // CUDA stream - auto stream = at::cuda::getCurrentCUDAStream(); + auto stream = getCurrentCUDAStream(); // Swizzle row-wise scales if needed - std::optional rowwise_scales_pyt; + std::optional rowwise_scales_pyt; if (rowwise_usage) { // Buffer for unswizzled scales const auto input_scales_nvte = tensor.get_rowwise_scale_inv(); @@ -102,7 +102,7 @@ std::tuple, std::optional> swizzle_scales_ } // Swizzle column-wise scales if needed - std::optional columnwise_scales_pyt; + std::optional columnwise_scales_pyt; if (columnwise_usage) { // Buffer for unswizzled scales const auto input_scales_nvte = tensor.get_columnwise_scale_inv(); @@ -143,7 +143,7 @@ std::tuple, std::optional> swizzle_scales_ namespace { -std::optional multi_tensor_swizzle_scales_for_gemm_impl( +std::optional multi_tensor_swizzle_scales_for_gemm_impl( std::vector &tensors, bool rowwise_usage, bool columnwise_usage, bool check_scale_inv_shapes) { // Checks and trivial cases @@ -260,10 +260,10 @@ std::optional multi_tensor_swizzle_scales_for_gemm_impl( NVTE_SCOPED_GIL_RELEASE({ if (check_scale_inv_shapes) { nvte_multi_tensor_swizzle_scaling_factors(inputs_nvte, outputs_nvte, n_swizzle, - at::cuda::getCurrentCUDAStream()); + getCurrentCUDAStream()); } else { nvte_multi_tensor_swizzle_scaling_factors_unchecked(inputs_nvte, outputs_nvte, n_swizzle, - at::cuda::getCurrentCUDAStream()); + getCurrentCUDAStream()); } }); @@ -286,21 +286,21 @@ std::optional multi_tensor_swizzle_scales_for_gemm_impl( } // anonymous namespace -std::optional multi_tensor_swizzle_scales_for_gemm( +std::optional multi_tensor_swizzle_scales_for_gemm( std::vector &tensors, bool rowwise_usage, bool columnwise_usage) { return multi_tensor_swizzle_scales_for_gemm_impl(tensors, rowwise_usage, columnwise_usage, /*check_scale_inv_shapes=*/true); } -std::optional multi_tensor_swizzle_scales_for_gemm_unchecked( +std::optional multi_tensor_swizzle_scales_for_gemm_unchecked( std::vector &tensors, bool rowwise_usage, bool columnwise_usage) { return multi_tensor_swizzle_scales_for_gemm_impl(tensors, rowwise_usage, columnwise_usage, /*check_scale_inv_shapes=*/false); } -at::Tensor convert_block_scaling_to_mxfp8_tensor(transformer_engine::TensorWrapper &input, +Tensor convert_block_scaling_to_mxfp8_tensor(transformer_engine::TensorWrapper &input, bool rowwise) { // Check input tensor const NVTEScalingMode scaling_mode = input.scaling_mode(); @@ -330,7 +330,7 @@ at::Tensor convert_block_scaling_to_mxfp8_tensor(transformer_engine::TensorWrapp const size_t swizzled_scale_inv_first_dim = ceildiv(data_flat_first_dim, 128) * 128; const size_t swizzled_scale_inv_last_dim = ceildiv(data_flat_last_dim, 128) * 4; // Allocate memory for swizzled mxfp8 scaling factors - at::Tensor swizzled_scale_inv = + Tensor swizzled_scale_inv = allocateSpace(std::vector{swizzled_scale_inv_first_dim, swizzled_scale_inv_last_dim}, transformer_engine::DType::kByte, false); // Set rowwise scaling factors on output @@ -346,7 +346,7 @@ at::Tensor convert_block_scaling_to_mxfp8_tensor(transformer_engine::TensorWrapp // Convert scaling factors from FP8 block scaling GEMM_READY format to mxfp8 swizzled format NVTE_SCOPED_GIL_RELEASE({ nvte_swizzle_block_scaling_to_mxfp8_scaling_factors(input_cu.data(), output_cu.data(), - at::cuda::getCurrentCUDAStream()); + getCurrentCUDAStream()); }); // Set the input tensor to be the converted mxfp8 tensor and return the swizzled scaling factor @@ -373,8 +373,8 @@ std::optional maybe_swizzle_grouped_tensor(GroupedTensorW return std::nullopt; } - std::optional rowwise_scales_pyt; - std::optional columnwise_scales_pyt; + std::optional rowwise_scales_pyt; + std::optional columnwise_scales_pyt; GroupedTensorWrapper swizzle_input(input.num_tensors(), input.logical_shape(), input.scaling_mode()); @@ -441,7 +441,7 @@ std::optional maybe_swizzle_grouped_tensor(GroupedTensorW NVTE_SCOPED_GIL_RELEASE({ nvte_swizzle_grouped_scaling_factors(swizzle_input.data(), swizzle_output.data(), - at::cuda::getCurrentCUDAStream()); + getCurrentCUDAStream()); }); if (swizzle_rowwise) { diff --git a/transformer_engine/pytorch/csrc/extensions/transpose.cpp b/transformer_engine/pytorch/csrc/extensions/transpose.cpp index 0318978195..498952337d 100644 --- a/transformer_engine/pytorch/csrc/extensions/transpose.cpp +++ b/transformer_engine/pytorch/csrc/extensions/transpose.cpp @@ -17,7 +17,7 @@ namespace transformer_engine { namespace pytorch { -at::Tensor fp8_transpose(at::Tensor input, DType otype, std::optional output) { +Tensor fp8_transpose(Tensor input, DType otype, std::optional output) { init_extension(); // Tensor dimensions @@ -32,12 +32,12 @@ at::Tensor fp8_transpose(at::Tensor input, DType otype, std::optional{M, N}, otype); auto output_cu = makeTransformerEngineTensor(out.data_ptr(), std::vector{N, M}, otype); - nvte_transpose(input_cu.data(), output_cu.data(), at::cuda::getCurrentCUDAStream()); + nvte_transpose(input_cu.data(), output_cu.data(), getCurrentCUDAStream()); return out; } -at::Tensor nvfp4_data_transpose(at::Tensor input, std::optional output) { +Tensor nvfp4_data_transpose(Tensor input, std::optional output) { init_extension(); // Input is packed FP4: logical [M, K] stored as [M, K/2] bytes @@ -72,15 +72,15 @@ at::Tensor nvfp4_data_transpose(at::Tensor input, std::optional outp std::vector output_shape = {static_cast(K), static_cast(M_packed)}; // Output tensor - at::Tensor out; + Tensor out; if (output.has_value()) { out = *output; NVTE_CHECK( static_cast(out.size(0)) == K && static_cast(out.size(1)) == M_packed, "Output shape mismatch for NVFP4 transpose."); } else { - const auto opts = at::TensorOptions().dtype(torch::kUInt8).device(torch::kCUDA); - out = at::empty(output_shape, opts); + const auto opts = TensorOptions().dtype(kUInt8).device(kCUDA); + out = empty(output_shape, opts); } // Return immediately if tensor is empty @@ -93,12 +93,12 @@ at::Tensor nvfp4_data_transpose(at::Tensor input, std::optional outp makeTransformerEngineTensor(input.data_ptr(), std::vector{M, K_packed}, DType::kByte); auto output_cu = makeTransformerEngineTensor(out.data_ptr(), std::vector{K, M_packed}, DType::kByte); - nvte_nvfp4_data_transpose(input_cu.data(), output_cu.data(), at::cuda::getCurrentCUDAStream()); + nvte_nvfp4_data_transpose(input_cu.data(), output_cu.data(), getCurrentCUDAStream()); return out; } -void nvfp4_2d_scale_transpose(at::Tensor input, at::Tensor output, int64_t M_tiles, +void nvfp4_2d_scale_transpose(Tensor input, Tensor output, int64_t M_tiles, int64_t K_tiles) { init_extension(); @@ -108,8 +108,8 @@ void nvfp4_2d_scale_transpose(at::Tensor input, at::Tensor output, int64_t M_til const auto out_shape = getTensorShape(output); NVTE_CHECK(in_shape.size() == 2, "NVFP4 scale transpose expects 2D input."); NVTE_CHECK(out_shape.size() == 2, "NVFP4 scale transpose expects 2D output."); - NVTE_CHECK(input.scalar_type() == at::kByte, "NVFP4 scale transpose input must be uint8 (E4M3)."); - NVTE_CHECK(output.scalar_type() == at::kByte, + NVTE_CHECK(input.scalar_type() == kByte, "NVFP4 scale transpose input must be uint8 (E4M3)."); + NVTE_CHECK(output.scalar_type() == kByte, "NVFP4 scale transpose output must be uint8 (E4M3)."); auto input_cu = makeTransformerEngineTensor( @@ -118,10 +118,10 @@ void nvfp4_2d_scale_transpose(at::Tensor input, at::Tensor output, int64_t M_til output.data_ptr(), std::vector{out_shape[0], out_shape[1]}, DType::kByte); nvte_nvfp4_scale_transpose(input_cu.data(), output_cu.data(), static_cast(M_tiles), - static_cast(K_tiles), at::cuda::getCurrentCUDAStream()); + static_cast(K_tiles), getCurrentCUDAStream()); } -void nvfp4_expand_scale_to_fp8(at::Tensor input, at::Tensor output, int64_t tile_rows, +void nvfp4_expand_scale_to_fp8(Tensor input, Tensor output, int64_t tile_rows, int64_t tile_cols, int64_t rows_padded, int64_t block_len) { init_extension(); @@ -131,8 +131,8 @@ void nvfp4_expand_scale_to_fp8(at::Tensor input, at::Tensor output, int64_t tile const auto out_shape = getTensorShape(output); NVTE_CHECK(in_shape.size() == 2, "NVFP4 expand scale expects 2D input."); NVTE_CHECK(out_shape.size() == 2, "NVFP4 expand scale expects 2D output."); - NVTE_CHECK(input.scalar_type() == at::kFloat, "NVFP4 expand scale input must be float32."); - NVTE_CHECK(output.scalar_type() == at::kByte, "NVFP4 expand scale output must be uint8 (E4M3)."); + NVTE_CHECK(input.scalar_type() == kFloat, "NVFP4 expand scale input must be float32."); + NVTE_CHECK(output.scalar_type() == kByte, "NVFP4 expand scale output must be uint8 (E4M3)."); auto input_cu = makeTransformerEngineTensor( input.data_ptr(), std::vector{in_shape[0], in_shape[1]}, DType::kFloat32); @@ -141,18 +141,18 @@ void nvfp4_expand_scale_to_fp8(at::Tensor input, at::Tensor output, int64_t tile nvte_nvfp4_expand_scale_to_fp8(input_cu.data(), output_cu.data(), static_cast(tile_rows), static_cast(tile_cols), static_cast(rows_padded), - static_cast(block_len), at::cuda::getCurrentCUDAStream()); + static_cast(block_len), getCurrentCUDAStream()); } -void nvfp4_compute_per_block_scale(at::Tensor block_amax, at::Tensor scale, - at::Tensor global_amax) { +void nvfp4_compute_per_block_scale(Tensor block_amax, Tensor scale, + Tensor global_amax) { init_extension(); // block_amax and scale: [tile_rows, tile_cols], float32 // global_amax: single element tensor, float32 (avoids D2H transfer) - NVTE_CHECK(block_amax.scalar_type() == at::kFloat, "Block amax must be float32."); - NVTE_CHECK(scale.scalar_type() == at::kFloat, "Scale must be float32."); - NVTE_CHECK(global_amax.scalar_type() == at::kFloat, "Global amax must be float32."); + NVTE_CHECK(block_amax.scalar_type() == kFloat, "Block amax must be float32."); + NVTE_CHECK(scale.scalar_type() == kFloat, "Scale must be float32."); + NVTE_CHECK(global_amax.scalar_type() == kFloat, "Global amax must be float32."); NVTE_CHECK(global_amax.numel() == 1, "Global amax must be a single element tensor."); auto block_amax_cu = makeTransformerEngineTensor(block_amax); @@ -160,11 +160,11 @@ void nvfp4_compute_per_block_scale(at::Tensor block_amax, at::Tensor scale, auto global_amax_cu = makeTransformerEngineTensor(global_amax); nvte_nvfp4_compute_per_block_scale(block_amax_cu.data(), scale_cu.data(), global_amax_cu.data(), - at::cuda::getCurrentCUDAStream()); + getCurrentCUDAStream()); } -void nvfp4_fused_scale(at::Tensor block_amax, at::Tensor global_amax, at::Tensor per_block_scale, - at::Tensor target_scale, at::Tensor target_amax, int64_t tile_rows, +void nvfp4_fused_scale(Tensor block_amax, Tensor global_amax, Tensor per_block_scale, + Tensor target_scale, Tensor target_amax, int64_t tile_rows, int64_t tile_cols, int64_t rows_padded, int64_t block_len) { init_extension(); @@ -173,11 +173,11 @@ void nvfp4_fused_scale(at::Tensor block_amax, at::Tensor global_amax, at::Tensor // per_block_scale: [tile_rows, tile_cols], float32 (for partial_cast) // target_scale: [rows_padded, tile_cols], uint8 (E4M3) // target_amax: [1], float32 - NVTE_CHECK(block_amax.scalar_type() == at::kFloat, "Block amax must be float32."); - NVTE_CHECK(global_amax.scalar_type() == at::kFloat, "Global amax must be float32."); - NVTE_CHECK(per_block_scale.scalar_type() == at::kFloat, "Per-block scale must be float32."); - NVTE_CHECK(target_scale.scalar_type() == at::kByte, "Target scale must be uint8 (E4M3)."); - NVTE_CHECK(target_amax.scalar_type() == at::kFloat, "Target amax must be float32."); + NVTE_CHECK(block_amax.scalar_type() == kFloat, "Block amax must be float32."); + NVTE_CHECK(global_amax.scalar_type() == kFloat, "Global amax must be float32."); + NVTE_CHECK(per_block_scale.scalar_type() == kFloat, "Per-block scale must be float32."); + NVTE_CHECK(target_scale.scalar_type() == kByte, "Target scale must be uint8 (E4M3)."); + NVTE_CHECK(target_amax.scalar_type() == kFloat, "Target amax must be float32."); NVTE_CHECK(global_amax.numel() == 1, "Global amax must be a single element tensor."); NVTE_CHECK(target_amax.numel() == 1, "Target amax must be a single element tensor."); @@ -191,13 +191,13 @@ void nvfp4_fused_scale(at::Tensor block_amax, at::Tensor global_amax, at::Tensor target_scale_cu.data(), target_amax_cu.data(), static_cast(tile_rows), static_cast(tile_cols), static_cast(rows_padded), static_cast(block_len), - at::cuda::getCurrentCUDAStream()); + getCurrentCUDAStream()); } void nvfp4_multi_tensor_fused_scale( - std::vector block_amax_list, std::vector global_amax_list, - std::vector per_block_scale_list, std::vector target_scale_list, - std::vector target_amax_list, std::vector tile_rows_list, + std::vector block_amax_list, std::vector global_amax_list, + std::vector per_block_scale_list, std::vector target_scale_list, + std::vector target_amax_list, std::vector tile_rows_list, std::vector tile_cols_list, std::vector rows_padded_list, int64_t block_len) { init_extension(); @@ -214,7 +214,7 @@ void nvfp4_multi_tensor_fused_scale( return; } - auto stream = at::cuda::getCurrentCUDAStream(); + auto stream = getCurrentCUDAStream(); for (size_t i = 0; i < num_tensors; ++i) { const auto& block_amax = block_amax_list[i]; @@ -226,11 +226,11 @@ void nvfp4_multi_tensor_fused_scale( const size_t tile_cols = static_cast(tile_cols_list[i]); const size_t rows_padded = static_cast(rows_padded_list[i]); - NVTE_CHECK(block_amax.scalar_type() == at::kFloat, "Block amax must be float32."); - NVTE_CHECK(global_amax.scalar_type() == at::kFloat, "Global amax must be float32."); - NVTE_CHECK(per_block_scale.scalar_type() == at::kFloat, "Per-block scale must be float32."); - NVTE_CHECK(target_scale.scalar_type() == at::kByte, "Target scale must be uint8 (E4M3)."); - NVTE_CHECK(target_amax.scalar_type() == at::kFloat, "Target amax must be float32."); + NVTE_CHECK(block_amax.scalar_type() == kFloat, "Block amax must be float32."); + NVTE_CHECK(global_amax.scalar_type() == kFloat, "Global amax must be float32."); + NVTE_CHECK(per_block_scale.scalar_type() == kFloat, "Per-block scale must be float32."); + NVTE_CHECK(target_scale.scalar_type() == kByte, "Target scale must be uint8 (E4M3)."); + NVTE_CHECK(target_amax.scalar_type() == kFloat, "Target amax must be float32."); NVTE_CHECK(global_amax.numel() == 1, "Global amax must be a single element tensor."); NVTE_CHECK(target_amax.numel() == 1, "Target amax must be a single element tensor."); @@ -246,21 +246,21 @@ void nvfp4_multi_tensor_fused_scale( } } -void nvfp4_compute_global_scale(at::Tensor global_amax, at::Tensor global_scale) { +void nvfp4_compute_global_scale(Tensor global_amax, Tensor global_scale) { init_extension(); // global_amax and global_scale: [num_params], float32 - NVTE_CHECK(global_amax.scalar_type() == at::kFloat, "Global amax must be float32."); - NVTE_CHECK(global_scale.scalar_type() == at::kFloat, "Global scale must be float32."); + NVTE_CHECK(global_amax.scalar_type() == kFloat, "Global amax must be float32."); + NVTE_CHECK(global_scale.scalar_type() == kFloat, "Global scale must be float32."); auto global_amax_cu = makeTransformerEngineTensor(global_amax); auto global_scale_cu = makeTransformerEngineTensor(global_scale); nvte_nvfp4_compute_global_scale(global_amax_cu.data(), global_scale_cu.data(), - at::cuda::getCurrentCUDAStream()); + getCurrentCUDAStream()); } -at::Tensor swap_first_dims(at::Tensor tensor, std::optional out) { +Tensor swap_first_dims(Tensor tensor, std::optional out) { init_extension(); // Make sure input is contiguous @@ -273,22 +273,22 @@ at::Tensor swap_first_dims(at::Tensor tensor, std::optional out) { std::vector out_shape_int64(in_shape.begin(), in_shape.end()); out_shape_int64[0] = static_cast(in_shape[1]); out_shape_int64[1] = static_cast(in_shape[0]); - auto opts = at::TensorOptions().dtype(input.dtype()).device(input.device()); - out = at::empty(out_shape_int64, opts); + auto opts = TensorOptions().dtype(input.dtype()).device(input.device()); + out = empty(out_shape_int64, opts); } // Launch kernel const TensorWrapper te_input = makeTransformerEngineTensor(input); TensorWrapper te_output = makeTransformerEngineTensor(*out); - nvte_swap_first_dims(te_input.data(), te_output.data(), at::cuda::getCurrentCUDAStream()); + nvte_swap_first_dims(te_input.data(), te_output.data(), getCurrentCUDAStream()); return std::move(*out); } -void nvfp4_2d_multi_tensor_transpose(std::vector rowwise_data_list, - std::vector columnwise_data_list, - std::vector rowwise_scale_inv_list, - std::vector columnwise_scale_inv_list, +void nvfp4_2d_multi_tensor_transpose(std::vector rowwise_data_list, + std::vector columnwise_data_list, + std::vector rowwise_scale_inv_list, + std::vector columnwise_scale_inv_list, std::vector M_list, std::vector K_list) { init_extension(); @@ -303,7 +303,7 @@ void nvfp4_2d_multi_tensor_transpose(std::vector rowwise_data_list, return; } - auto stream = at::cuda::getCurrentCUDAStream(); + auto stream = getCurrentCUDAStream(); // Process each tensor - the main benefit is reduced Python overhead // by doing the iteration in C++ rather than Python diff --git a/transformer_engine/pytorch/csrc/pybind.h b/transformer_engine/pytorch/csrc/pybind.h index 9e640537f9..67100d5a65 100644 --- a/transformer_engine/pytorch/csrc/pybind.h +++ b/transformer_engine/pytorch/csrc/pybind.h @@ -13,7 +13,6 @@ #include #include #include -#include #include "common.h" #include "transformer_engine/transformer_engine.h" @@ -99,8 +98,8 @@ TensorWrapper NVTETensorFromNVFP4Tensor(py::handle tensor, Quantizer *quantizer) GroupedTensorWrapper GroupedTensorFromPyTorchGroupedTensor(py::handle tensor); -inline bool IsFloatingPointType(at::ScalarType type) { - return type == at::kFloat || type == at::kHalf || type == at::kBFloat16; +inline bool IsFloatingPointType(ScalarType type) { + return type == kFloat || type == kHalf || type == kBFloat16; } constexpr std::array custom_types_converters = { diff --git a/transformer_engine/pytorch/csrc/quantizer.cpp b/transformer_engine/pytorch/csrc/quantizer.cpp index 5fc50953a1..5efa8b45f3 100644 --- a/transformer_engine/pytorch/csrc/quantizer.cpp +++ b/transformer_engine/pytorch/csrc/quantizer.cpp @@ -10,7 +10,6 @@ #include "common/util/cuda_runtime.h" #include "common/util/system.h" #include "pybind.h" -#include "torch/torch.h" namespace transformer_engine::pytorch { @@ -20,8 +19,8 @@ namespace { * * If no device is provided, uses the current CUDA device. */ -at::Device resolve_device(std::optional device, - const std::optional& data = std::nullopt) { +Device resolve_device(std::optional device, + const std::optional& data = std::nullopt) { if (device.has_value() && data.has_value()) { // Ensure that they are the same const auto provided_device = *device; @@ -36,7 +35,7 @@ at::Device resolve_device(std::optional device, if (data.has_value()) { return data->device(); } - return at::Device(torch::kCUDA, c10::cuda::current_device()); + return Device(kCUDA, current_device()); } /*! @brief Transposed tensor shape @@ -84,8 +83,8 @@ std::vector convert_shape_for_fp4(const std::vector& shape) { return ret; } -std::optional build_grouped_tensor_offsets(const size_t num_tensors, - const std::optional& first_dims, +std::optional build_grouped_tensor_offsets(const size_t num_tensors, + const std::optional& first_dims, const size_t logical_last_dim) { if (!first_dims.has_value()) { return std::nullopt; @@ -93,27 +92,27 @@ std::optional build_grouped_tensor_offsets(const size_t num_tensors, const auto& first_dims_tensor = first_dims.value(); NVTE_CHECK(first_dims_tensor.is_cuda(), "first_dims must be on CUDA."); - NVTE_CHECK(first_dims_tensor.scalar_type() == at::kLong, "first_dims must have dtype int64."); + NVTE_CHECK(first_dims_tensor.scalar_type() == kLong, "first_dims must have dtype int64."); NVTE_CHECK(static_cast(first_dims_tensor.numel()) == num_tensors, "first_dims must have length ", num_tensors, "."); const int64_t logical_last_dim_i64 = static_cast(logical_last_dim); const auto first_dims_contiguous = first_dims_tensor.contiguous(); auto tensor_offsets = - at::empty({static_cast(num_tensors) + 1}, first_dims_contiguous.options()); + empty({static_cast(num_tensors) + 1}, first_dims_contiguous.options()); NVTE_SCOPED_GIL_RELEASE({ nvte_splits_to_offsets(static_cast(first_dims_contiguous.data_ptr()), static_cast(tensor_offsets.data_ptr()), num_tensors, - logical_last_dim_i64, at::cuda::getCurrentCUDAStream()); + logical_last_dim_i64, getCurrentCUDAStream()); }); return tensor_offsets; } -at::TensorOptions grouped_tensor_data_options(const DType dtype) { - return at::TensorOptions().dtype(GetATenDType(dtype)).device(torch::kCUDA); +TensorOptions grouped_tensor_data_options(const DType dtype) { + return TensorOptions().dtype(GetATenDType(dtype)).device(kCUDA); } -py::object maybe_tensor_to_py(const std::optional& tensor) { +py::object maybe_tensor_to_py(const std::optional& tensor) { return tensor ? py::cast(*tensor) : py::none(); } @@ -143,8 +142,8 @@ Quantizer::Quantizer(const py::handle& quantizer) { } Float8Quantizer::Float8Quantizer(const py::handle& quantizer) : Quantizer(quantizer) { - const at::Tensor& scale = quantizer.attr("scale").cast(); - const at::Tensor& amax = quantizer.attr("amax").cast(); + const Tensor& scale = quantizer.attr("scale").cast(); + const Tensor& amax = quantizer.attr("amax").cast(); const DType type = quantizer.attr("dtype").cast(); this->amax = amax; @@ -153,18 +152,18 @@ Float8Quantizer::Float8Quantizer(const py::handle& quantizer) : Quantizer(quanti } std::pair NoneQuantizer::create_tensor( - const std::vector& shape, DType dtype, std::optional device_opt, + const std::vector& shape, DType dtype, std::optional device_opt, bool pin_memory) const { const auto device = resolve_device(device_opt); const std::vector shape_int64(shape.begin(), shape.end()); const auto opts = - at::TensorOptions().dtype(GetATenDType(dtype)).device(device).pinned_memory(pin_memory); - return create_tensor(shape, dtype, at::empty(shape_int64, opts)); + TensorOptions().dtype(GetATenDType(dtype)).device(device).pinned_memory(pin_memory); + return create_tensor(shape, dtype, empty(shape_int64, opts)); } std::pair NoneQuantizer::create_tensor(const std::vector& shape, DType dtype, - at::Tensor data) const { + Tensor data) const { TensorWrapper out_cpp; out_cpp.set_rowwise_data(data.data_ptr(), dtype, shape); set_quantization_params(&out_cpp); @@ -173,8 +172,8 @@ std::pair NoneQuantizer::create_tensor(const std::vec std::pair NoneQuantizer::create_grouped_tensor( const size_t num_tensors, const std::vector& logical_shape, const DType dtype, - py::object quantizer, const std::optional& first_dims, - const std::optional& precomputed_tensor_offsets, const size_t logical_first_dim, + py::object quantizer, const std::optional& first_dims, + const std::optional& precomputed_tensor_offsets, const size_t logical_first_dim, const size_t logical_last_dim) const { using namespace pybind11::literals; @@ -185,15 +184,15 @@ std::pair NoneQuantizer::create_grouped_tensor const int64_t total_elements = static_cast(logical_first_dim) * static_cast(logical_last_dim); - std::optional rowwise_data; - std::optional columnwise_data; + std::optional rowwise_data; + std::optional columnwise_data; const bool with_rowwise_data = rowwise_usage; const bool with_columnwise_data = columnwise_usage; if (with_rowwise_data) { - rowwise_data = at::empty({total_elements}, grouped_tensor_data_options(dtype)); + rowwise_data = empty({total_elements}, grouped_tensor_data_options(dtype)); } if (with_columnwise_data) { - columnwise_data = at::empty({total_elements}, grouped_tensor_data_options(dtype)); + columnwise_data = empty({total_elements}, grouped_tensor_data_options(dtype)); } GroupedTensorWrapper out_cpp(num_tensors, logical_shape, this->get_scaling_mode()); @@ -246,7 +245,7 @@ std::pair NoneQuantizer::create_grouped_tensor std::pair NoneQuantizer::convert_and_update_tensor( py::object tensor) const { - auto tensor_pyt = tensor.cast(); + auto tensor_pyt = tensor.cast(); TensorWrapper out_cpp; out_cpp.set_rowwise_data(tensor_pyt.data_ptr(), GetTransformerEngineDType(tensor_pyt.scalar_type()), @@ -263,26 +262,26 @@ void NoneQuantizer::quantize(const TensorWrapper& input, TensorWrapper& out, void Float8Quantizer::set_quantization_params(TensorWrapper* tensor) const { tensor->set_scale(scale.data_ptr(), GetTransformerEngineDType(scale.scalar_type()), getTensorShape(scale)); - at::TensorOptions opts = opts.dtype(torch::kFloat32).device(torch::kCUDA); + TensorOptions opts = opts.dtype(kFloat32).device(kCUDA); tensor->set_amax(amax.data_ptr(), GetTransformerEngineDType(amax.scalar_type()), getTensorShape(amax)); } std::pair Float8Quantizer::create_tensor( - const std::vector& shape, DType dtype, std::optional device_opt, + const std::vector& shape, DType dtype, std::optional device_opt, bool pin_memory) const { const auto device = resolve_device(device_opt); const auto opts = - at::TensorOptions().dtype(torch::kFloat32).device(device).pinned_memory(pin_memory); - at::Tensor scale_inv = at::empty(std::vector{1}, opts); + TensorOptions().dtype(kFloat32).device(device).pinned_memory(pin_memory); + Tensor scale_inv = empty(std::vector{1}, opts); return create_tensor(shape, dtype, std::nullopt, std::nullopt, std::move(scale_inv), device, pin_memory); } std::pair Float8Quantizer::create_tensor( - const std::vector& shape, DType dtype, std::optional data, - std::optional transpose, std::optional scale_inv, - std::optional device_opt, bool pin_memory) const { + const std::vector& shape, DType dtype, std::optional data, + std::optional transpose, std::optional scale_inv, + std::optional device_opt, bool pin_memory) const { const auto device = resolve_device(device_opt, data); using namespace pybind11::literals; int is_non_tn_fp8_gemm_supported = nvte_is_non_tn_fp8_gemm_supported(); @@ -291,8 +290,8 @@ std::pair Float8Quantizer::create_tensor( if (with_data && !data) { const std::vector shape_int64(shape.begin(), shape.end()); const auto opts = - at::TensorOptions().dtype(torch::kUInt8).device(device).pinned_memory(pin_memory); - data = at::empty(shape_int64, opts); + TensorOptions().dtype(kUInt8).device(device).pinned_memory(pin_memory); + data = empty(shape_int64, opts); } else if (!with_data && data) { data.reset(); } @@ -303,15 +302,15 @@ std::pair Float8Quantizer::create_tensor( if (with_transpose && !transpose) { const auto transpose_shape = make_transpose_shape(shape); const auto opts = - at::TensorOptions().dtype(torch::kUInt8).device(device).pinned_memory(pin_memory); - transpose = at::empty(transpose_shape, opts); + TensorOptions().dtype(kUInt8).device(device).pinned_memory(pin_memory); + transpose = empty(transpose_shape, opts); } else if (!with_transpose && transpose) { transpose.reset(); } py::object transpose_py = with_transpose ? py::cast(*transpose) : py::none(); // Initialize scale-inverse tensor if (!scale_inv) { - scale_inv = at::reciprocal(scale); + scale_inv = reciprocal(scale); } py::object scale_inv_py = py::cast(*scale_inv); // Construct Python FP8 tensor @@ -379,8 +378,8 @@ std::pair Float8Quantizer::create_tensor( std::pair Float8Quantizer::create_grouped_tensor( const size_t num_tensors, const std::vector& logical_shape, const DType dtype, - py::object quantizer, const std::optional& first_dims, - const std::optional& precomputed_tensor_offsets, const size_t logical_first_dim, + py::object quantizer, const std::optional& first_dims, + const std::optional& precomputed_tensor_offsets, const size_t logical_first_dim, const size_t logical_last_dim) const { using namespace pybind11::literals; @@ -391,22 +390,22 @@ std::pair Float8Quantizer::create_grouped_tens const int64_t total_elements = static_cast(logical_first_dim) * static_cast(logical_last_dim); - const auto uint8_opts = at::TensorOptions().dtype(torch::kUInt8).device(torch::kCUDA); - const auto float_opts = at::TensorOptions().dtype(torch::kFloat32).device(torch::kCUDA); + const auto uint8_opts = TensorOptions().dtype(kUInt8).device(kCUDA); + const auto float_opts = TensorOptions().dtype(kFloat32).device(kCUDA); - std::optional rowwise_data; - std::optional columnwise_data; - std::optional rowwise_scale_inv; - std::optional columnwise_scale_inv; - at::Tensor amax = at::empty({static_cast(num_tensors)}, float_opts); + std::optional rowwise_data; + std::optional columnwise_data; + std::optional rowwise_scale_inv; + std::optional columnwise_scale_inv; + Tensor amax = empty({static_cast(num_tensors)}, float_opts); if (rowwise_usage) { - rowwise_data = at::empty({total_elements}, uint8_opts); - rowwise_scale_inv = at::empty({static_cast(num_tensors)}, float_opts); + rowwise_data = empty({total_elements}, uint8_opts); + rowwise_scale_inv = empty({static_cast(num_tensors)}, float_opts); } if (columnwise_usage) { - columnwise_data = at::empty({total_elements}, uint8_opts); - columnwise_scale_inv = at::empty({static_cast(num_tensors)}, float_opts); + columnwise_data = empty({total_elements}, uint8_opts); + columnwise_scale_inv = empty({static_cast(num_tensors)}, float_opts); } GroupedTensorWrapper out_cpp(num_tensors, logical_shape, this->get_scaling_mode()); @@ -477,14 +476,14 @@ std::pair Float8Quantizer::convert_and_update_tensor( const bool has_data = !data_py.is_none(); const bool has_transpose = !transpose_py.is_none(); NVTE_CHECK(has_data || has_transpose, "Float8Tensor has no data."); - std::optional data_tensor, transpose_tensor; + std::optional data_tensor, transpose_tensor; if (has_data) { - data_tensor = data_py.cast(); + data_tensor = data_py.cast(); } if (has_transpose) { - transpose_tensor = transpose_py.cast(); + transpose_tensor = transpose_py.cast(); } - at::Tensor scale_inv_tensor = tensor.attr("_scale_inv").cast(); + Tensor scale_inv_tensor = tensor.attr("_scale_inv").cast(); // Tensor dimensions std::vector shape; @@ -512,8 +511,8 @@ std::pair Float8Quantizer::convert_and_update_tensor( tensor.attr("_data") = data_py; } else if (!has_data && need_data) { const std::vector shape_int64(shape.begin(), shape.end()); - const auto opts = at::TensorOptions().dtype(torch::kUInt8).device(torch::kCUDA); - data_tensor = at::empty(shape_int64, opts); + const auto opts = TensorOptions().dtype(kUInt8).device(kCUDA); + data_tensor = empty(shape_int64, opts); data_py = py::cast(data_tensor); tensor.attr("_data") = data_py; } @@ -525,8 +524,8 @@ std::pair Float8Quantizer::convert_and_update_tensor( tensor.attr("_transpose") = transpose_py; } else if (!has_transpose && need_transpose) { const auto transpose_shape = make_transpose_shape(shape); - const auto opts = at::TensorOptions().dtype(torch::kUInt8).device(torch::kCUDA); - transpose_tensor = at::empty(transpose_shape, opts); + const auto opts = TensorOptions().dtype(kUInt8).device(kCUDA); + transpose_tensor = empty(transpose_shape, opts); transpose_py = py::cast(transpose_tensor); tensor.attr("_transpose") = transpose_py; } @@ -563,7 +562,7 @@ void Float8Quantizer::quantize(const TensorWrapper& input, TensorWrapper& out, quant_config.set_noop_tensor(noop_flag->data()); } NVTE_SCOPED_GIL_RELEASE({ - nvte_quantize_v2(input.data(), out.data(), quant_config, at::cuda::getCurrentCUDAStream()); + nvte_quantize_v2(input.data(), out.data(), quant_config, getCurrentCUDAStream()); }); } @@ -573,12 +572,12 @@ Float8CurrentScalingQuantizer::Float8CurrentScalingQuantizer(const py::handle& q // Get amax reduction group if needed const bool with_amax_reduction = quantizer.attr("with_amax_reduction").cast(); - c10::intrusive_ptr amax_reduction_group; + IntrusivePtr amax_reduction_group; if (with_amax_reduction) { auto group = quantizer.attr("_canonicalized_amax_reduction_group")(); NVTE_CHECK(!group.is_none(), "Float8CurrentScalingQuantizer could not canonicalize amax reduction group"); - amax_reduction_group = group.cast>(); + amax_reduction_group = group.cast>(); } this->with_amax_reduction = with_amax_reduction; this->amax_reduction_group = amax_reduction_group; @@ -591,38 +590,38 @@ Float8CurrentScalingQuantizer::Float8CurrentScalingQuantizer(const py::handle& q void Float8CurrentScalingQuantizer::set_quantization_params(TensorWrapper* tensor) const {} std::pair Float8CurrentScalingQuantizer::create_tensor( - const std::vector& shape, DType dtype, std::optional device_opt, + const std::vector& shape, DType dtype, std::optional device_opt, bool pin_memory) const { const auto device = resolve_device(device_opt); using namespace pybind11::literals; // Initialize data tensor - at::Tensor data_tensor; + Tensor data_tensor; int is_non_tn_fp8_gemm_supported = nvte_is_non_tn_fp8_gemm_supported(); const bool with_data = rowwise_usage || is_non_tn_fp8_gemm_supported; if (with_data) { const std::vector shape_int64(shape.begin(), shape.end()); const auto opts = - at::TensorOptions().dtype(torch::kUInt8).device(device).pinned_memory(pin_memory); - data_tensor = at::empty(shape_int64, opts); + TensorOptions().dtype(kUInt8).device(device).pinned_memory(pin_memory); + data_tensor = empty(shape_int64, opts); } // Initialize transpose tensor - at::Tensor transpose_tensor; + Tensor transpose_tensor; const bool with_transpose = columnwise_usage && !is_non_tn_fp8_gemm_supported; if (with_transpose) { const auto transpose_shape = make_transpose_shape(shape); const auto opts = - at::TensorOptions().dtype(torch::kUInt8).device(device).pinned_memory(pin_memory); - transpose_tensor = at::empty(transpose_shape, opts); + TensorOptions().dtype(kUInt8).device(device).pinned_memory(pin_memory); + transpose_tensor = empty(transpose_shape, opts); } // Initialize scale-inverse tensor - at::Tensor scale_inv_tensor; + Tensor scale_inv_tensor; { const std::vector scale_inv_shape = {1}; const auto opts = - at::TensorOptions().dtype(torch::kFloat32).device(device).pinned_memory(pin_memory); - scale_inv_tensor = at::empty(scale_inv_shape, opts); + TensorOptions().dtype(kFloat32).device(device).pinned_memory(pin_memory); + scale_inv_tensor = empty(scale_inv_shape, opts); } // Construct Python FP8 tensor py::object out_py; @@ -692,8 +691,8 @@ std::pair Float8CurrentScalingQuantizer::create_tenso std::pair Float8CurrentScalingQuantizer::create_grouped_tensor( const size_t num_tensors, const std::vector& logical_shape, const DType dtype, - py::object quantizer, const std::optional& first_dims, - const std::optional& precomputed_tensor_offsets, const size_t logical_first_dim, + py::object quantizer, const std::optional& first_dims, + const std::optional& precomputed_tensor_offsets, const size_t logical_first_dim, const size_t logical_last_dim) const { using namespace pybind11::literals; @@ -704,23 +703,23 @@ std::pair Float8CurrentScalingQuantizer::creat const int64_t total_elements = static_cast(logical_first_dim) * static_cast(logical_last_dim); - const auto uint8_opts = at::TensorOptions().dtype(torch::kUInt8).device(torch::kCUDA); - const auto float_opts = at::TensorOptions().dtype(torch::kFloat32).device(torch::kCUDA); + const auto uint8_opts = TensorOptions().dtype(kUInt8).device(kCUDA); + const auto float_opts = TensorOptions().dtype(kFloat32).device(kCUDA); - std::optional rowwise_data; - std::optional columnwise_data; - std::optional rowwise_scale_inv; - std::optional columnwise_scale_inv; - at::Tensor scale = at::empty({static_cast(num_tensors)}, float_opts); - at::Tensor amax = at::empty({static_cast(num_tensors)}, float_opts); + std::optional rowwise_data; + std::optional columnwise_data; + std::optional rowwise_scale_inv; + std::optional columnwise_scale_inv; + Tensor scale = empty({static_cast(num_tensors)}, float_opts); + Tensor amax = empty({static_cast(num_tensors)}, float_opts); if (rowwise_usage) { - rowwise_data = at::empty({total_elements}, uint8_opts); - rowwise_scale_inv = at::empty({static_cast(num_tensors)}, float_opts); + rowwise_data = empty({total_elements}, uint8_opts); + rowwise_scale_inv = empty({static_cast(num_tensors)}, float_opts); } if (columnwise_usage) { - columnwise_data = at::empty({total_elements}, uint8_opts); - columnwise_scale_inv = at::empty({static_cast(num_tensors)}, float_opts); + columnwise_data = empty({total_elements}, uint8_opts); + columnwise_scale_inv = empty({static_cast(num_tensors)}, float_opts); } GroupedTensorWrapper out_cpp(num_tensors, logical_shape, this->get_scaling_mode()); @@ -777,12 +776,12 @@ std::pair Float8CurrentScalingQuantizer::creat return {std::move(out_cpp), std::move(out_py)}; } -std::tuple +std::tuple Float8CurrentScalingQuantizer::create_unquantized_tensor_with_amax(const std::vector& shape, DType dtype, - std::optional data) { - const auto opts = at::TensorOptions().dtype(torch::kFloat32).device(torch::kCUDA); - at::Tensor amax_buf = at::zeros({1}, opts); + std::optional data) { + const auto opts = TensorOptions().dtype(kFloat32).device(kCUDA); + Tensor amax_buf = zeros({1}, opts); auto out = data.has_value() ? NoneQuantizer(py::none()).create_tensor(shape, dtype, data.value()) : NoneQuantizer(py::none()).create_tensor(shape, dtype); TensorWrapper out_cpp = std::move(out.first); @@ -807,14 +806,14 @@ std::pair Float8CurrentScalingQuantizer::convert_and_ const bool has_data = !data_py.is_none(); const bool has_transpose = !transpose_py.is_none(); NVTE_CHECK(has_data || has_transpose, "Tensor has no data."); - std::optional data_tensor, transpose_tensor; + std::optional data_tensor, transpose_tensor; if (has_data) { - data_tensor = data_py.cast(); + data_tensor = data_py.cast(); } if (has_transpose) { - transpose_tensor = transpose_py.cast(); + transpose_tensor = transpose_py.cast(); } - at::Tensor scale_inv_tensor = tensor.attr("_scale_inv").cast(); + Tensor scale_inv_tensor = tensor.attr("_scale_inv").cast(); // Tensor dimensions std::vector shape; @@ -842,8 +841,8 @@ std::pair Float8CurrentScalingQuantizer::convert_and_ tensor.attr("_data") = data_py; } else if (!has_data && need_data) { const std::vector shape_int64(shape.begin(), shape.end()); - const auto opts = at::TensorOptions().dtype(torch::kUInt8).device(torch::kCUDA); - data_tensor = at::empty(shape_int64, opts); + const auto opts = TensorOptions().dtype(kUInt8).device(kCUDA); + data_tensor = empty(shape_int64, opts); data_py = py::cast(data_tensor); tensor.attr("_data") = data_py; } @@ -855,8 +854,8 @@ std::pair Float8CurrentScalingQuantizer::convert_and_ tensor.attr("_transpose") = transpose_py; } else if (!has_transpose && need_transpose) { const auto transpose_shape = make_transpose_shape(shape); - const auto opts = at::TensorOptions().dtype(torch::kUInt8).device(torch::kCUDA); - transpose_tensor = at::empty(transpose_shape, opts); + const auto opts = TensorOptions().dtype(kUInt8).device(kCUDA); + transpose_tensor = empty(transpose_shape, opts); transpose_py = py::cast(transpose_tensor); tensor.attr("_transpose") = transpose_py; } @@ -885,12 +884,12 @@ std::pair Float8CurrentScalingQuantizer::convert_and_ void Float8CurrentScalingQuantizer::quantize_impl(const TensorWrapper& input, TensorWrapper& out, const std::optional& noop_flag, - bool compute_amax, at::Tensor amax_buf, - at::Tensor scale_buf) { + bool compute_amax, Tensor amax_buf, + Tensor scale_buf) { out.set_amax(amax_buf.data_ptr(), DType::kFloat32, std::vector{1}); out.set_scale(scale_buf.data_ptr(), DType::kFloat32, std::vector{1}); - auto stream = at::cuda::getCurrentCUDAStream(); + auto stream = getCurrentCUDAStream(); // Nothing to be done if input is empty if (input.numel() == 0) { @@ -919,9 +918,9 @@ void Float8CurrentScalingQuantizer::quantize_impl(const TensorWrapper& input, Te // Perform amax reduction if needed if (with_amax_reduction) { // allreduce amax tensor - c10d::AllreduceOptions opts; - opts.reduceOp = c10d::ReduceOp::MAX; - std::vector tensors = {amax_buf}; + AllreduceOptions opts; + opts.reduceOp = ReduceOp::MAX; + std::vector tensors = {amax_buf}; NVTE_SCOPED_GIL_RELEASE({ amax_reduction_group->allreduce(tensors, opts)->wait(); }); } @@ -939,17 +938,17 @@ void Float8CurrentScalingQuantizer::quantize_impl(const TensorWrapper& input, Te void Float8CurrentScalingQuantizer::quantize(const TensorWrapper& input, TensorWrapper& out, const std::optional& noop_flag) { - const auto opts = at::TensorOptions().dtype(torch::kFloat32).device(torch::kCUDA); - at::Tensor amax_and_scale = at::empty({2}, opts); + const auto opts = TensorOptions().dtype(kFloat32).device(kCUDA); + Tensor amax_and_scale = empty({2}, opts); this->quantize_impl(input, out, noop_flag, true, amax_and_scale[0], amax_and_scale[1]); } void Float8CurrentScalingQuantizer::quantize_with_amax( - TensorWrapper& input, TensorWrapper& out, at::Tensor amax, + TensorWrapper& input, TensorWrapper& out, Tensor amax, const std::optional& noop_flag) { - const auto opts = at::TensorOptions().dtype(torch::kFloat32).device(torch::kCUDA); + const auto opts = TensorOptions().dtype(kFloat32).device(kCUDA); input.set_amax(nullptr, DType::kFloat32, input.defaultShape); - this->quantize_impl(input, out, noop_flag, false, std::move(amax), at::empty({1}, opts)); + this->quantize_impl(input, out, noop_flag, false, std::move(amax), empty({1}, opts)); } Float8BlockQuantizer::Float8BlockQuantizer(const py::handle& quantizer) : Quantizer(quantizer) { @@ -964,7 +963,7 @@ Float8BlockQuantizer::Float8BlockQuantizer(const py::handle& quantizer) : Quanti void Float8BlockQuantizer::set_quantization_params(TensorWrapper* tensor) const {} std::pair Float8BlockQuantizer::create_tensor( - const std::vector& shape, DType dtype, std::optional device_opt, + const std::vector& shape, DType dtype, std::optional device_opt, bool pin_memory) const { const auto device = resolve_device(device_opt); using namespace pybind11::literals; @@ -974,19 +973,19 @@ std::pair Float8BlockQuantizer::create_tensor( } TensorWrapper tensor(this->get_scaling_mode()); - at::TensorOptions opts; - at::TensorOptions scale_opts; - at::Tensor data_rowwise, data_colwise, scale_inv_rowwise, scale_inv_colwise; - opts = opts.dtype(torch::kUInt8).device(device).pinned_memory(pin_memory); - scale_opts = scale_opts.dtype(torch::kFloat32).device(device).pinned_memory(pin_memory); + TensorOptions opts; + TensorOptions scale_opts; + Tensor data_rowwise, data_colwise, scale_inv_rowwise, scale_inv_colwise; + opts = opts.dtype(kUInt8).device(device).pinned_memory(pin_memory); + scale_opts = scale_opts.dtype(kFloat32).device(device).pinned_memory(pin_memory); if (rowwise_usage) { - data_rowwise = at::empty(torch_shape, opts); + data_rowwise = empty(torch_shape, opts); auto scale_shape = get_scale_shape(shape, false); size_t sinv0 = scale_shape[0]; size_t sinv1 = scale_shape[1]; scale_inv_rowwise = - at::empty({static_cast(sinv0), static_cast(sinv1)}, scale_opts); + empty({static_cast(sinv0), static_cast(sinv1)}, scale_opts); tensor.set_rowwise_data(data_rowwise.data_ptr(), this->dtype, shape); tensor.set_rowwise_scale_inv(scale_inv_rowwise.data_ptr(), DType::kFloat32, std::vector{sinv0, sinv1}); @@ -1010,9 +1009,9 @@ std::pair Float8BlockQuantizer::create_tensor( auto scale_shape = get_scale_shape(shape, true); size_t sinv0 = scale_shape[0]; size_t sinv1 = scale_shape[1]; - data_colwise = at::empty(torch_columnwise_shape, opts); + data_colwise = empty(torch_columnwise_shape, opts); scale_inv_colwise = - at::empty({static_cast(sinv0), static_cast(sinv1)}, scale_opts); + empty({static_cast(sinv0), static_cast(sinv1)}, scale_opts); tensor.set_columnwise_data(data_colwise.data_ptr(), this->dtype, columnwise_shape); tensor.set_columnwise_scale_inv(scale_inv_colwise.data_ptr(), DType::kFloat32, @@ -1074,8 +1073,8 @@ std::pair Float8BlockQuantizer::create_tensor( std::pair Float8BlockQuantizer::create_grouped_tensor( const size_t num_tensors, const std::vector& logical_shape, const DType dtype, - py::object quantizer, const std::optional& first_dims, - const std::optional& precomputed_tensor_offsets, const size_t logical_first_dim, + py::object quantizer, const std::optional& first_dims, + const std::optional& precomputed_tensor_offsets, const size_t logical_first_dim, const size_t logical_last_dim) const { using namespace pybind11::literals; @@ -1086,27 +1085,27 @@ std::pair Float8BlockQuantizer::create_grouped const int64_t total_elements = static_cast(logical_first_dim) * static_cast(logical_last_dim); - const auto uint8_opts = at::TensorOptions().dtype(torch::kUInt8).device(torch::kCUDA); - const auto float_opts = at::TensorOptions().dtype(torch::kFloat32).device(torch::kCUDA); + const auto uint8_opts = TensorOptions().dtype(kUInt8).device(kCUDA); + const auto float_opts = TensorOptions().dtype(kFloat32).device(kCUDA); - std::optional rowwise_data; - std::optional columnwise_data; - std::optional rowwise_scale_inv; - std::optional columnwise_scale_inv; + std::optional rowwise_data; + std::optional columnwise_data; + std::optional rowwise_scale_inv; + std::optional columnwise_scale_inv; const std::vector logical_shape_vec = {logical_first_dim, logical_last_dim}; if (rowwise_usage) { - rowwise_data = at::empty({total_elements}, uint8_opts); + rowwise_data = empty({total_elements}, uint8_opts); const auto scale_shape = get_scale_shape(logical_shape_vec, false); const int64_t total_scale_elements = static_cast(product(scale_shape)); - rowwise_scale_inv = at::empty({total_scale_elements}, float_opts); + rowwise_scale_inv = empty({total_scale_elements}, float_opts); } if (columnwise_usage) { - columnwise_data = at::empty({total_elements}, uint8_opts); + columnwise_data = empty({total_elements}, uint8_opts); const auto scale_shape = get_scale_shape(logical_shape_vec, true); const int64_t total_scale_elements = static_cast(product(scale_shape)); - columnwise_scale_inv = at::empty({total_scale_elements}, float_opts); + columnwise_scale_inv = empty({total_scale_elements}, float_opts); } GroupedTensorWrapper out_cpp(num_tensors, logical_shape, this->get_scaling_mode()); @@ -1168,12 +1167,12 @@ std::pair Float8BlockQuantizer::convert_and_update_te const bool with_gemm_swizzled_scales = true; // Extract buffers from Python tensor - auto get_tensor = [&tensor](const char* name) -> std::optional { + auto get_tensor = [&tensor](const char* name) -> std::optional { auto attr_py = tensor.attr(name); if (attr_py.is_none()) { return std::nullopt; } - return attr_py.cast(); + return attr_py.cast(); }; auto rowwise_data = get_tensor("_rowwise_data"); auto rowwise_scale_inv = get_tensor("_rowwise_scale_inv"); @@ -1182,10 +1181,10 @@ std::pair Float8BlockQuantizer::convert_and_update_te NVTE_CHECK(rowwise_data || columnwise_data, "FP8BlockwiseTensor has no data."); // Tensor options and dimensions - at::TensorOptions opts; - at::TensorOptions scale_opts; - opts = opts.dtype(torch::kUInt8).device(torch::kCUDA); - scale_opts = scale_opts.dtype(torch::kFloat32).device(torch::kCUDA); + TensorOptions opts; + TensorOptions scale_opts; + opts = opts.dtype(kUInt8).device(kCUDA); + scale_opts = scale_opts.dtype(kFloat32).device(kCUDA); auto get_columnwise_shape = [&columnwise_data]() -> std::vector { if (!columnwise_data) { @@ -1220,7 +1219,7 @@ std::pair Float8BlockQuantizer::convert_and_update_te // Coerce row-wise data if (rowwise_usage) { if (!rowwise_data) { - rowwise_data = at::empty(torch_shape, opts); + rowwise_data = empty(torch_shape, opts); tensor.attr("_rowwise_data") = *rowwise_data; } if (!rowwise_scale_inv) { @@ -1228,7 +1227,7 @@ std::pair Float8BlockQuantizer::convert_and_update_te size_t sinv0 = scale_shape[0]; size_t sinv1 = scale_shape[1]; rowwise_scale_inv = - at::empty({static_cast(sinv0), static_cast(sinv1)}, scale_opts); + empty({static_cast(sinv0), static_cast(sinv1)}, scale_opts); tensor.attr("_rowwise_scale_inv") = *rowwise_scale_inv; } } else { // rowwise_usage == false @@ -1257,7 +1256,7 @@ std::pair Float8BlockQuantizer::convert_and_update_te } } if (!columnwise_data) { - columnwise_data = at::empty(torch_columnwise_shape, opts); + columnwise_data = empty(torch_columnwise_shape, opts); tensor.attr("_columnwise_data") = *columnwise_data; } if (!columnwise_scale_inv) { @@ -1265,7 +1264,7 @@ std::pair Float8BlockQuantizer::convert_and_update_te size_t sinv0 = scale_shape[0]; size_t sinv1 = scale_shape[1]; columnwise_scale_inv = - at::empty({static_cast(sinv0), static_cast(sinv1)}, scale_opts); + empty({static_cast(sinv0), static_cast(sinv1)}, scale_opts); tensor.attr("_columnwise_scale_inv") = *columnwise_scale_inv; } } else { // columnwise_usage == false @@ -1282,8 +1281,8 @@ std::pair Float8BlockQuantizer::convert_and_update_te auto ret = TensorWrapper(is_2D_scaled ? NVTE_BLOCK_SCALING_2D : NVTE_BLOCK_SCALING_1D); if (rowwise_usage) { - const at::Tensor& data_rowwise = tensor.attr("_rowwise_data").cast(); - const at::Tensor& scale_inv_rowwise = tensor.attr("_rowwise_scale_inv").cast(); + const Tensor& data_rowwise = tensor.attr("_rowwise_data").cast(); + const Tensor& scale_inv_rowwise = tensor.attr("_rowwise_scale_inv").cast(); void* scale_inv_rowwise_dptr = scale_inv_rowwise.data_ptr(); const auto& rowwise_shape = getTensorShape(data_rowwise); ret.set_rowwise_data(data_rowwise.data_ptr(), dtype, rowwise_shape); @@ -1291,8 +1290,8 @@ std::pair Float8BlockQuantizer::convert_and_update_te ret.set_rowwise_scale_inv(scale_inv_rowwise_dptr, DType::kFloat32, scale_inv_rowwise_shape); } if (columnwise_usage) { - const at::Tensor& data_colwise = tensor.attr("_columnwise_data").cast(); - const at::Tensor& scale_inv_colwise = tensor.attr("_columnwise_scale_inv").cast(); + const Tensor& data_colwise = tensor.attr("_columnwise_data").cast(); + const Tensor& scale_inv_colwise = tensor.attr("_columnwise_scale_inv").cast(); void* scale_inv_colwise_dptr = scale_inv_colwise.data_ptr(); const auto& shape = getTensorShape(data_colwise); ret.set_columnwise_data(data_colwise.data_ptr(), dtype, shape); @@ -1316,7 +1315,7 @@ void Float8BlockQuantizer::quantize(const TensorWrapper& input, TensorWrapper& o quant_config.set_force_pow_2_scales(force_pow_2_scales); quant_config.set_amax_epsilon(amax_epsilon); NVTE_SCOPED_GIL_RELEASE({ - nvte_quantize_v2(input.data(), out.data(), quant_config, at::cuda::getCurrentCUDAStream()); + nvte_quantize_v2(input.data(), out.data(), quant_config, getCurrentCUDAStream()); }); } @@ -1381,7 +1380,7 @@ MXFP8Quantizer::MXFP8Quantizer(const py::handle& quantizer) : Quantizer(quantize void MXFP8Quantizer::set_quantization_params(TensorWrapper* tensor) const {} std::pair MXFP8Quantizer::create_tensor( - const std::vector& shape, DType dtype, std::optional device_opt, + const std::vector& shape, DType dtype, std::optional device_opt, bool pin_memory) const { const auto device = resolve_device(device_opt); using namespace pybind11::literals; @@ -1399,25 +1398,25 @@ std::pair MXFP8Quantizer::create_tensor( const auto columnwise_scale_inv_shape = get_scale_shape(shape, true); // Allocate tensors - at::Tensor rowwise_data_tensor, rowwise_scale_inv_tensor; - at::Tensor columnwise_data_tensor, columnwise_scale_inv_tensor; + Tensor rowwise_data_tensor, rowwise_scale_inv_tensor; + Tensor columnwise_data_tensor, columnwise_scale_inv_tensor; const auto uint8_tensor_opts = - at::TensorOptions().dtype(torch::kUInt8).device(device).pinned_memory(pin_memory); + TensorOptions().dtype(kUInt8).device(device).pinned_memory(pin_memory); if (rowwise_usage) { const std::vector scale_inv_shape_int64(rowwise_scale_inv_shape.begin(), rowwise_scale_inv_shape.end()); - rowwise_data_tensor = at::empty(shape_int64, uint8_tensor_opts); - rowwise_scale_inv_tensor = at::empty(scale_inv_shape_int64, uint8_tensor_opts); + rowwise_data_tensor = empty(shape_int64, uint8_tensor_opts); + rowwise_scale_inv_tensor = empty(scale_inv_shape_int64, uint8_tensor_opts); } if (columnwise_usage) { const std::vector scale_inv_shape_int64(columnwise_scale_inv_shape.begin(), columnwise_scale_inv_shape.end()); - columnwise_data_tensor = at::empty(shape_int64, uint8_tensor_opts); - columnwise_scale_inv_tensor = at::empty(scale_inv_shape_int64, uint8_tensor_opts); + columnwise_data_tensor = empty(shape_int64, uint8_tensor_opts); + columnwise_scale_inv_tensor = empty(scale_inv_shape_int64, uint8_tensor_opts); } // Convert tensors to Python - auto py_cast = [](at::Tensor& tensor, bool need_cast) -> py::object { + auto py_cast = [](Tensor& tensor, bool need_cast) -> py::object { return need_cast ? py::cast(tensor) : py::none(); }; auto rowwise_data_py = py_cast(rowwise_data_tensor, rowwise_usage); @@ -1495,8 +1494,8 @@ std::pair MXFP8Quantizer::create_tensor( std::pair MXFP8Quantizer::create_grouped_tensor( const size_t num_tensors, const std::vector& logical_shape, const DType dtype, - py::object quantizer, const std::optional& first_dims, - const std::optional& precomputed_tensor_offsets, const size_t logical_first_dim, + py::object quantizer, const std::optional& first_dims, + const std::optional& precomputed_tensor_offsets, const size_t logical_first_dim, const size_t logical_last_dim) const { using namespace pybind11::literals; @@ -1507,26 +1506,26 @@ std::pair MXFP8Quantizer::create_grouped_tenso const int64_t total_elements = static_cast(logical_first_dim) * static_cast(logical_last_dim); - const auto uint8_opts = at::TensorOptions().dtype(torch::kUInt8).device(torch::kCUDA); + const auto uint8_opts = TensorOptions().dtype(kUInt8).device(kCUDA); - std::optional rowwise_data; - std::optional columnwise_data; - std::optional rowwise_scale_inv; - std::optional columnwise_scale_inv; + std::optional rowwise_data; + std::optional columnwise_data; + std::optional rowwise_scale_inv; + std::optional columnwise_scale_inv; const std::vector logical_shape_vec = {logical_first_dim, logical_last_dim}; if (rowwise_usage) { - rowwise_data = at::empty({total_elements}, uint8_opts); + rowwise_data = empty({total_elements}, uint8_opts); const auto scale_shape = get_scale_shape(logical_shape_vec, false); const int64_t total_scale_elements = static_cast(product(scale_shape)); - rowwise_scale_inv = at::empty({total_scale_elements}, uint8_opts); + rowwise_scale_inv = empty({total_scale_elements}, uint8_opts); } if (columnwise_usage) { - columnwise_data = at::empty({total_elements}, uint8_opts); + columnwise_data = empty({total_elements}, uint8_opts); const auto scale_shape = get_scale_shape(logical_shape_vec, true); const int64_t total_scale_elements = static_cast(product(scale_shape)); - columnwise_scale_inv = at::empty({total_scale_elements}, uint8_opts); + columnwise_scale_inv = empty({total_scale_elements}, uint8_opts); } GroupedTensorWrapper out_cpp(num_tensors, logical_shape, this->get_scaling_mode()); @@ -1591,12 +1590,12 @@ std::pair MXFP8Quantizer::convert_and_update_tensor( const bool with_gemm_swizzled_scales = this->optimize_for_gemm; // Extract buffers from Python tensor - auto get_tensor = [&tensor](const char* name) -> std::optional { + auto get_tensor = [&tensor](const char* name) -> std::optional { auto attr_py = tensor.attr(name); if (attr_py.is_none()) { return std::nullopt; } - return attr_py.cast(); + return attr_py.cast(); }; auto rowwise_data = get_tensor("_rowwise_data"); auto rowwise_scale_inv = get_tensor("_rowwise_scale_inv"); @@ -1621,16 +1620,16 @@ std::pair MXFP8Quantizer::convert_and_update_tensor( if (rowwise_usage) { if (!rowwise_data) { const std::vector shape_int64(shape.begin(), shape.end()); - const auto opts = at::TensorOptions().dtype(torch::kUInt8).device(torch::kCUDA); - rowwise_data = at::empty(shape_int64, opts); + const auto opts = TensorOptions().dtype(kUInt8).device(kCUDA); + rowwise_data = empty(shape_int64, opts); tensor.attr("_rowwise_data") = *rowwise_data; } if (!rowwise_scale_inv) { const auto scale_inv_shape = get_scale_shape(shape, false); const std::vector scale_inv_shape_int64(scale_inv_shape.begin(), scale_inv_shape.end()); - const auto opts = at::TensorOptions().dtype(torch::kUInt8).device(torch::kCUDA); - rowwise_scale_inv = at::empty(scale_inv_shape_int64, opts); + const auto opts = TensorOptions().dtype(kUInt8).device(kCUDA); + rowwise_scale_inv = empty(scale_inv_shape_int64, opts); tensor.attr("_rowwise_scale_inv") = *rowwise_scale_inv; } } else { // rowwise_usage == false @@ -1648,16 +1647,16 @@ std::pair MXFP8Quantizer::convert_and_update_tensor( if (columnwise_usage) { if (!columnwise_data) { const std::vector shape_int64(shape.begin(), shape.end()); - const auto opts = at::TensorOptions().dtype(torch::kUInt8).device(torch::kCUDA); - columnwise_data = at::empty(shape_int64, opts); + const auto opts = TensorOptions().dtype(kUInt8).device(kCUDA); + columnwise_data = empty(shape_int64, opts); tensor.attr("_columnwise_data") = *columnwise_data; } if (!columnwise_scale_inv) { const auto scale_inv_shape = get_scale_shape(shape, true); const std::vector scale_inv_shape_int64(scale_inv_shape.begin(), scale_inv_shape.end()); - const auto opts = at::TensorOptions().dtype(torch::kUInt8).device(torch::kCUDA); - columnwise_scale_inv = at::empty(scale_inv_shape_int64, opts); + const auto opts = TensorOptions().dtype(kUInt8).device(kCUDA); + columnwise_scale_inv = empty(scale_inv_shape_int64, opts); tensor.attr("_columnwise_scale_inv") = *columnwise_scale_inv; } } else { // columnwise_usage == false @@ -1703,7 +1702,7 @@ void MXFP8Quantizer::quantize(const TensorWrapper& input, TensorWrapper& out, quant_config.set_noop_tensor(noop_flag->data()); } NVTE_SCOPED_GIL_RELEASE({ - nvte_quantize_v2(input.data(), out.data(), quant_config, at::cuda::getCurrentCUDAStream()); + nvte_quantize_v2(input.data(), out.data(), quant_config, getCurrentCUDAStream()); }); } @@ -1762,17 +1761,17 @@ NVFP4Quantizer::NVFP4Quantizer(const py::handle& quantizer) : Quantizer(quantize // Get amax reduction group if needed for NVFP4 AG const bool with_amax_reduction = quantizer.attr("with_amax_reduction").cast(); - c10::intrusive_ptr amax_reduction_group; + IntrusivePtr amax_reduction_group; if (with_amax_reduction) { auto group = quantizer.attr("_canonicalized_amax_reduction_group")(); NVTE_CHECK(!group.is_none(), "NVFP4Quantizer could not canonicalize amax reduction group"); - amax_reduction_group = group.cast>(); + amax_reduction_group = group.cast>(); } this->with_amax_reduction = with_amax_reduction; this->amax_reduction_group = amax_reduction_group; this->rht_matrix_random_sign_mask_t = quantizer.attr("rht_matrix_random_sign_mask_t").cast(); - this->rht_matrix = quantizer.attr("rht_matrix").cast(); + this->rht_matrix = quantizer.attr("rht_matrix").cast(); } void NVFP4Quantizer::set_quantization_params(TensorWrapper* tensor) const { @@ -1798,7 +1797,7 @@ bool NVFP4Quantizer::is_eligible_for_rht_cast_fusion(const std::vector& } std::pair NVFP4Quantizer::create_tensor( - const std::vector& shape, DType dtype, std::optional device_opt, + const std::vector& shape, DType dtype, std::optional device_opt, bool pin_memory) const { const auto device = resolve_device(device_opt); using namespace pybind11::literals; @@ -1828,21 +1827,21 @@ std::pair NVFP4Quantizer::create_tensor( const auto columnwise_scale_inv_shape = get_scale_shape(shape, true); // Allocate tensors - at::Tensor rowwise_data_tensor, rowwise_scale_inv_tensor, amax_rowwise; - at::Tensor columnwise_data_tensor, columnwise_scale_inv_tensor, amax_columnwise; + Tensor rowwise_data_tensor, rowwise_scale_inv_tensor, amax_rowwise; + Tensor columnwise_data_tensor, columnwise_scale_inv_tensor, amax_columnwise; const auto bit8_tensor_opts = - at::TensorOptions().dtype(torch::kUInt8).device(device).pinned_memory(pin_memory); + TensorOptions().dtype(kUInt8).device(device).pinned_memory(pin_memory); const auto bit32_tensor_opts = - at::TensorOptions().dtype(torch::kFloat32).device(device).pinned_memory(pin_memory); + TensorOptions().dtype(kFloat32).device(device).pinned_memory(pin_memory); if (rowwise_usage) { const std::vector scale_inv_shape_int64(rowwise_scale_inv_shape.begin(), rowwise_scale_inv_shape.end()); - rowwise_data_tensor = at::empty(convert_shape_for_fp4(shape_int64), bit8_tensor_opts); - rowwise_scale_inv_tensor = at::empty(scale_inv_shape_int64, bit8_tensor_opts); + rowwise_data_tensor = empty(convert_shape_for_fp4(shape_int64), bit8_tensor_opts); + rowwise_scale_inv_tensor = empty(scale_inv_shape_int64, bit8_tensor_opts); const int64_t amax_rows = row_scaled_nvfp4 ? static_cast(flat_first_dim) : 1; // hadamard amax kernel will zero out pointer with ZeroAmaxKernel // nvte_compute_amax_with_config will zero out the pointer if needed - amax_rowwise = at::empty({amax_rows}, bit32_tensor_opts); + amax_rowwise = empty({amax_rows}, bit32_tensor_opts); } if (columnwise_usage) { const std::vector scale_inv_shape_int64(columnwise_scale_inv_shape.begin(), @@ -1853,15 +1852,15 @@ std::pair NVFP4Quantizer::create_tensor( static_cast(flat_last_dim)}; const auto transpose_shape_int64 = make_transpose_shape(shape_int64_2d); columnwise_data_tensor = - at::empty(convert_shape_for_fp4(transpose_shape_int64), bit8_tensor_opts); - columnwise_scale_inv_tensor = at::empty(scale_inv_shape_int64, bit8_tensor_opts); + empty(convert_shape_for_fp4(transpose_shape_int64), bit8_tensor_opts); + columnwise_scale_inv_tensor = empty(scale_inv_shape_int64, bit8_tensor_opts); // hadamard amax kernel will zero out pointer with ZeroAmaxKernel // nvte_compute_amax_with_config will zero out the pointer if needed - amax_columnwise = at::empty({1}, bit32_tensor_opts); + amax_columnwise = empty({1}, bit32_tensor_opts); } // Convert tensors to Python - auto py_cast = [](at::Tensor& tensor, bool need_cast) -> py::object { + auto py_cast = [](Tensor& tensor, bool need_cast) -> py::object { return need_cast ? py::cast(tensor) : py::none(); }; auto rowwise_data_py = py_cast(rowwise_data_tensor, rowwise_usage); @@ -1961,8 +1960,8 @@ std::pair NVFP4Quantizer::create_tensor( std::pair NVFP4Quantizer::create_grouped_tensor( const size_t num_tensors, const std::vector& logical_shape, const DType dtype, - py::object quantizer, const std::optional& first_dims, - const std::optional& precomputed_tensor_offsets, const size_t logical_first_dim, + py::object quantizer, const std::optional& first_dims, + const std::optional& precomputed_tensor_offsets, const size_t logical_first_dim, const size_t logical_last_dim) const { using namespace pybind11::literals; @@ -1974,15 +1973,15 @@ std::pair NVFP4Quantizer::create_grouped_tenso static_cast(logical_first_dim) * static_cast(logical_last_dim); NVTE_CHECK(total_elements % 2 == 0, "NVFP4 data size must be divisible by 2."); - const auto uint8_opts = at::TensorOptions().dtype(torch::kUInt8).device(torch::kCUDA); - const auto float_opts = at::TensorOptions().dtype(torch::kFloat32).device(torch::kCUDA); + const auto uint8_opts = TensorOptions().dtype(kUInt8).device(kCUDA); + const auto float_opts = TensorOptions().dtype(kFloat32).device(kCUDA); - std::optional rowwise_data; - std::optional columnwise_data; - std::optional rowwise_scale_inv; - std::optional columnwise_scale_inv; - std::optional rowwise_amax; - std::optional columnwise_amax; + std::optional rowwise_data; + std::optional columnwise_data; + std::optional rowwise_scale_inv; + std::optional columnwise_scale_inv; + std::optional rowwise_amax; + std::optional columnwise_amax; const std::vector logical_shape_vec = {logical_first_dim, logical_last_dim}; const bool row_scaled_nvfp4 = this->row_scaled_nvfp4; const bool nvfp4_use_4over6 = this->nvfp4_4over6_mode != kNVTENVFP44Over6Disabled; @@ -1996,21 +1995,21 @@ std::pair NVFP4Quantizer::create_grouped_tenso const int64_t total_data_elements = total_elements / 2; if (rowwise_usage) { - rowwise_data = at::empty({total_data_elements}, uint8_opts); + rowwise_data = empty({total_data_elements}, uint8_opts); const auto scale_shape = get_scale_shape(logical_shape_vec, false); const int64_t total_scale_elements = static_cast(product(scale_shape)); - rowwise_scale_inv = at::empty({total_scale_elements}, uint8_opts); + rowwise_scale_inv = empty({total_scale_elements}, uint8_opts); const int64_t amax_elements = row_scaled_nvfp4 ? static_cast(logical_first_dim) : static_cast(num_tensors); - rowwise_amax = at::empty({amax_elements}, float_opts); + rowwise_amax = empty({amax_elements}, float_opts); } if (columnwise_usage) { - columnwise_data = at::empty({total_data_elements}, uint8_opts); + columnwise_data = empty({total_data_elements}, uint8_opts); const auto scale_shape = get_scale_shape(logical_shape_vec, true); const int64_t total_scale_elements = static_cast(product(scale_shape)); - columnwise_scale_inv = at::empty({total_scale_elements}, uint8_opts); - columnwise_amax = at::empty({static_cast(num_tensors)}, float_opts); + columnwise_scale_inv = empty({total_scale_elements}, uint8_opts); + columnwise_amax = empty({static_cast(num_tensors)}, float_opts); } GroupedTensorWrapper out_cpp(num_tensors, logical_shape, this->get_scaling_mode()); @@ -2095,7 +2094,7 @@ std::pair NVFP4Quantizer::create_unquantized_tensor_w // Zero out amax const size_t amax_numel = product(amax_shape); NVTE_CHECK_CUDA( - cudaMemsetAsync(amax_ptr, 0, amax_numel * sizeof(float), at::cuda::getCurrentCUDAStream())); + cudaMemsetAsync(amax_ptr, 0, amax_numel * sizeof(float), getCurrentCUDAStream())); return {std::move(out_cpp), std::move(out_py)}; } @@ -2105,12 +2104,12 @@ std::pair NVFP4Quantizer::convert_and_update_tensor( NVTE_CHECK(detail::IsNVFP4Tensor(tensor.ptr()), "NVFP4Quantizer must output to IsNVFP4Tensor."); // Extract buffers from Python tensor - auto get_tensor = [&tensor](const char* name) -> std::optional { + auto get_tensor = [&tensor](const char* name) -> std::optional { auto attr_py = tensor.attr(name); if (attr_py.is_none()) { return std::nullopt; } - return attr_py.cast(); + return attr_py.cast(); }; auto rowwise_data = get_tensor("_rowwise_data"); auto rowwise_scale_inv = get_tensor("_rowwise_scale_inv"); @@ -2157,24 +2156,24 @@ std::pair NVFP4Quantizer::convert_and_update_tensor( if (rowwise_usage) { if (!rowwise_data) { const std::vector shape_int64(shape.begin(), shape.end()); - const auto opts = at::TensorOptions().dtype(torch::kUInt8).device(torch::kCUDA); - rowwise_data = at::empty(convert_shape_for_fp4(shape_int64), opts); + const auto opts = TensorOptions().dtype(kUInt8).device(kCUDA); + rowwise_data = empty(convert_shape_for_fp4(shape_int64), opts); tensor.attr("_rowwise_data") = *rowwise_data; } if (!rowwise_scale_inv) { const auto scale_inv_shape = get_scale_shape(shape, false); const std::vector scale_inv_shape_int64(scale_inv_shape.begin(), scale_inv_shape.end()); - const auto opts = at::TensorOptions().dtype(torch::kUInt8).device(torch::kCUDA); - rowwise_scale_inv = at::empty(scale_inv_shape_int64, opts); + const auto opts = TensorOptions().dtype(kUInt8).device(kCUDA); + rowwise_scale_inv = empty(scale_inv_shape_int64, opts); tensor.attr("_rowwise_scale_inv") = *rowwise_scale_inv; } const int64_t amax_rows = row_scaled_nvfp4 ? static_cast(flat_first_dim) : 1; if (!amax_rowwise || amax_rowwise->numel() != amax_rows) { - const auto opts = at::TensorOptions().dtype(torch::kFloat32).device(torch::kCUDA); + const auto opts = TensorOptions().dtype(kFloat32).device(kCUDA); // hadamard amax kernel will zero out pointer with ZeroAmaxKernel // nvte_compute_amax_with_config will zero out the pointer if needed - amax_rowwise = at::empty({amax_rows}, opts); + amax_rowwise = empty({amax_rows}, opts); tensor.attr("_amax_rowwise") = *amax_rowwise; } } else { // rowwise_usage == false @@ -2199,24 +2198,24 @@ std::pair NVFP4Quantizer::convert_and_update_tensor( // and the transposed shape is [H, S, B], so divide last dim by 2 gives zero std::vector shape_int64_2d = {static_cast(flat_first_dim), static_cast(flat_last_dim)}; - const auto opts = at::TensorOptions().dtype(torch::kUInt8).device(torch::kCUDA); + const auto opts = TensorOptions().dtype(kUInt8).device(kCUDA); const auto transpose_shape_int64 = make_transpose_shape(shape_int64_2d); - columnwise_data = at::empty(convert_shape_for_fp4(transpose_shape_int64), opts); + columnwise_data = empty(convert_shape_for_fp4(transpose_shape_int64), opts); tensor.attr("_columnwise_data") = *columnwise_data; } if (!columnwise_scale_inv) { const auto scale_inv_shape = get_scale_shape(shape, true); const std::vector scale_inv_shape_int64(scale_inv_shape.begin(), scale_inv_shape.end()); - const auto opts = at::TensorOptions().dtype(torch::kUInt8).device(torch::kCUDA); - columnwise_scale_inv = at::empty(scale_inv_shape_int64, opts); + const auto opts = TensorOptions().dtype(kUInt8).device(kCUDA); + columnwise_scale_inv = empty(scale_inv_shape_int64, opts); tensor.attr("_columnwise_scale_inv") = *columnwise_scale_inv; } if (!amax_columnwise) { - const auto opts = at::TensorOptions().dtype(torch::kFloat32).device(torch::kCUDA); + const auto opts = TensorOptions().dtype(kFloat32).device(kCUDA); // hadamard amax kernel will zero out pointer with ZeroAmaxKernel // nvte_compute_amax_with_config will zero out the pointer if needed - amax_columnwise = at::empty({1}, opts); + amax_columnwise = empty({1}, opts); tensor.attr("_amax_columnwise") = *amax_columnwise; } } else { // columnwise_usage == false @@ -2344,13 +2343,13 @@ void NVFP4Quantizer::quantize_impl(const TensorWrapper& input, TensorWrapper& ou return; } - std::vector amax_tensors; + std::vector amax_tensors; auto make_amax_tensor = [](void* data_ptr) { NVTE_CHECK(data_ptr != nullptr, "Could not find amax pointer for NVFP4 amax reduction."); - return at::from_blob( + return from_blob( data_ptr, std::vector{1}, [](void*) {}, // deleter doing nothing since it doesn't own the data - at::device(at::kCUDA).dtype(torch::kFloat32)); + TensorOptions().device(kCUDA).dtype(kFloat32)); }; if (rowwise_usage) { amax_tensors.push_back(make_amax_tensor(out.get_amax().data_ptr)); @@ -2362,8 +2361,8 @@ void NVFP4Quantizer::quantize_impl(const TensorWrapper& input, TensorWrapper& ou return; } - c10d::AllreduceCoalescedOptions opts; - opts.reduceOp = c10d::ReduceOp::MAX; + AllreduceCoalescedOptions opts; + opts.reduceOp = ReduceOp::MAX; NVTE_SCOPED_GIL_RELEASE( { this->amax_reduction_group->allreduce_coalesced(amax_tensors, opts)->wait(); }); }; @@ -2376,7 +2375,7 @@ void NVFP4Quantizer::quantize_impl(const TensorWrapper& input, TensorWrapper& ou return; } - auto stream = at::cuda::getCurrentCUDAStream(); + auto stream = getCurrentCUDAStream(); QuantizationConfigWrapper quant_config; QuantizationConfigWrapper quant_config_columnwise; @@ -2434,21 +2433,21 @@ void NVFP4Quantizer::quantize_impl(const TensorWrapper& input, TensorWrapper& ou if (this->stochastic_rounding) { const size_t rng_elts_per_thread = 1024; // Wild guess, probably can be tightened - auto gen = at::get_generator_or_default( - std::nullopt, at::cuda::detail::getDefaultCUDAGenerator()); - auto opts = at::TensorOptions().dtype(torch::kInt64).device(torch::kCUDA); + auto gen = get_generator_or_default( + std::nullopt, getDefaultCUDAGenerator()); + auto opts = TensorOptions().dtype(kInt64).device(kCUDA); // Generate RNG state for rowwise quantization - at::PhiloxCudaState philox_args = init_philox_state(gen, rng_elts_per_thread); - auto rng_state = torch::empty({2}, opts); + PhiloxCudaState philox_args = init_philox_state(gen, rng_elts_per_thread); + auto rng_state = empty({2}, opts); philox_unpack(philox_args, static_cast(rng_state.data_ptr())); te_rng_state = makeTransformerEngineTensor(rng_state); quant_config.set_rng_state(te_rng_state.data()); // Generate separate RNG state for columnwise quantization if (need_separate_columnwise_rng) { - at::PhiloxCudaState philox_args_columnwise = init_philox_state(gen, rng_elts_per_thread); - auto rng_state_columnwise = torch::empty({2}, opts); + PhiloxCudaState philox_args_columnwise = init_philox_state(gen, rng_elts_per_thread); + auto rng_state_columnwise = empty({2}, opts); philox_unpack(philox_args_columnwise, static_cast(rng_state_columnwise.data_ptr())); te_rng_state_columnwise = makeTransformerEngineTensor(rng_state_columnwise); quant_config_columnwise.set_stochastic_rounding(true); @@ -2547,7 +2546,7 @@ void NVFP4Quantizer::quantize_impl(const TensorWrapper& input, TensorWrapper& ou auto& columnwise_quant_config_to_use = need_separate_columnwise_rng ? quant_config_columnwise : quant_config; // unfused path also needs memory allocation for intermediate buffer for RHT output - at::Tensor rht_output_t; // The RHT(x_t) output, in columnwise layout + Tensor rht_output_t; // The RHT(x_t) output, in columnwise layout // This wrapper is going to be passed as input to the quantization kernel. TensorWrapper rht_output_t_cpp; // Wrapper to contain the RHT(x) and RHT(x_t) outputs rht_output_t = @@ -2582,12 +2581,12 @@ void NVFP4Quantizer::quantize_with_amax(TensorWrapper& input, TensorWrapper& out if (input_amax_ptr != output_rowwise_amax_ptr && input_amax_ptr != nullptr && output_rowwise_amax_ptr != nullptr) { NVTE_CHECK_CUDA(cudaMemcpyAsync(output_rowwise_amax_ptr, input_amax_ptr, sizeof(float), - cudaMemcpyDeviceToDevice, at::cuda::getCurrentCUDAStream())); + cudaMemcpyDeviceToDevice, getCurrentCUDAStream())); } if (input_amax_ptr != output_columnwise_amax_ptr && input_amax_ptr != nullptr && output_columnwise_amax_ptr != nullptr) { NVTE_CHECK_CUDA(cudaMemcpyAsync(output_columnwise_amax_ptr, input_amax_ptr, sizeof(float), - cudaMemcpyDeviceToDevice, at::cuda::getCurrentCUDAStream())); + cudaMemcpyDeviceToDevice, getCurrentCUDAStream())); } input.set_amax(nullptr, DType::kFloat32, input.defaultShape); diff --git a/transformer_engine/pytorch/csrc/torch_backend.cpp b/transformer_engine/pytorch/csrc/torch_backend.cpp new file mode 100644 index 0000000000..39d86da406 --- /dev/null +++ b/transformer_engine/pytorch/csrc/torch_backend.cpp @@ -0,0 +1,16 @@ +/************************************************************************* + * Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * + * See LICENSE for license information. + ************************************************************************/ + +// Translation unit for the single TE<->PyTorch binary boundary. +// +// The facade helpers (GetATenDType, GetTransformerEngineDType, new_cuda_tensor) +// are defined `inline` in torch_backend.h: they sit on the per-tensor +// Python<->C++ marshalling path, so keeping them inlinable matches the +// pre-facade codegen. This .cpp is intentionally empty; it exists as the home +// for any future non-inline boundary code (and, together with torch_backend.h, +// is the only place allowed to name at::/c10::/torch:: symbols). + +#include "torch_backend.h" diff --git a/transformer_engine/pytorch/csrc/torch_backend.h b/transformer_engine/pytorch/csrc/torch_backend.h new file mode 100644 index 0000000000..848104be84 --- /dev/null +++ b/transformer_engine/pytorch/csrc/torch_backend.h @@ -0,0 +1,246 @@ +/************************************************************************* + * Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * + * See LICENSE for license information. + ************************************************************************/ + +/*! \file torch_backend.h + * \brief Single binary boundary between Transformer Engine and PyTorch. + * + * This header is the ONLY place in the PyTorch extension that is allowed to + * include libtorch/ATen/c10 headers and to name ``at::`` / ``c10::`` / + * ``torch::`` symbols. Every other translation unit talks to PyTorch + * exclusively through the type aliases (``using``) and free functions + * (``methods``) declared here. + * + * Why: it lets us swap the concrete tensor implementation from the classic + * ``at::Tensor`` to ``torch::stable::Tensor`` (the LibTorch stable ABI) by + * flipping a single compile flag, without touching the ~40 extension files. + * It also collapses the TE<->torch ABI surface to one auditable file, which a + * lint guard (``qa/L0_pytorch_lint/check_torch_boundary.sh``) enforces. + * + * Switch: + * - default -> classic libtorch (``at::Tensor``) + * - -DTE_WITH_STABLE_ABI -> torch stable ABI (``torch::stable::Tensor``) + */ + +#ifndef TRANSFORMER_ENGINE_PYTORCH_CSRC_TORCH_BACKEND_H_ +#define TRANSFORMER_ENGINE_PYTORCH_CSRC_TORCH_BACKEND_H_ + +// =========================================================================== +// The one and only libtorch include site. +// =========================================================================== +#ifdef TE_WITH_STABLE_ABI +// --------------------------------------------------------------------------- +// Stable-ABI path (migration target). torch::stable::Tensor exposes a narrower +// interface than at::Tensor, so operations are routed through the free-function +// wrappers below (implemented in torch_backend.cpp against the stable C shim). +// The aliases that have a direct stable counterpart are provided here; anything +// still missing a stable wrapper is tracked in the branch CLAUDE.md. +// --------------------------------------------------------------------------- +#include +#include +#else +// --------------------------------------------------------------------------- +// Classic libtorch path (default). +// --------------------------------------------------------------------------- +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +#include "c10/util/ArrayRef.h" +#endif + +#include +#include + +#include +#include +#include + +#include "common/util/logging.h" // NVTE_ERROR used by the inline helpers below + +namespace transformer_engine::pytorch { + +// =========================================================================== +// Type aliases -- the "usings" every extension file must use instead of naming +// at::/c10::/torch:: types directly. +// =========================================================================== +#ifdef TE_WITH_STABLE_ABI + +using Tensor = torch::stable::Tensor; +// NOTE: ScalarType/Device/Stream/ProcessGroup/PhiloxCudaState/CUDAGeneratorImpl +// stable equivalents are wired incrementally; see branch CLAUDE.md. + +#else + +// --- Types ----------------------------------------------------------------- +using Tensor = at::Tensor; // also torch::Tensor +using ScalarType = at::ScalarType; // also c10::ScalarType +using Device = at::Device; // also c10::Device +using Stream = at::Stream; +using TensorOptions = at::TensorOptions; // also torch::TensorOptions +using IntArrayRef = c10::IntArrayRef; +template +using ArrayRef = c10::ArrayRef; +template +using IntrusivePtr = c10::intrusive_ptr; +using Generator = at::Generator; +using CustomClassHolder = torch::CustomClassHolder; + +// Distributed process group (python: torch.distributed.ProcessGroup) and the +// collective option structs used for amax reduction. +using ProcessGroup = c10d::ProcessGroup; +using ReduceOp = c10d::ReduceOp; +using c10d::AllreduceCoalescedOptions; +using c10d::AllreduceOptions; +using c10d::BroadcastOptions; + +// CUDA interop. +using PhiloxCudaState = at::PhiloxCudaState; +using CUDAGeneratorImpl = at::CUDAGeneratorImpl; +using CUDAStream = at::cuda::CUDAStream; +using CUDAGuard = at::cuda::CUDAGuard; + +// --- dtype / device constants (torch-free spellings) ----------------------- +// RHS keeps the original torch spelling on purpose (this is the boundary). +inline constexpr auto kCUDA = at::kCUDA; +inline constexpr auto kCPU = at::kCPU; +inline constexpr auto kByte = at::kByte; +inline constexpr auto kUInt8 = torch::kUInt8; +inline constexpr auto kInt8 = torch::kInt8; +inline constexpr auto kInt32 = torch::kInt32; +inline constexpr auto kInt64 = torch::kInt64; +inline constexpr auto kLong = at::kLong; +inline constexpr auto kFloat = at::kFloat; +inline constexpr auto kFloat32 = torch::kFloat32; +inline constexpr auto kHalf = at::kHalf; +inline constexpr auto kBFloat16 = at::kBFloat16; +inline constexpr auto kBool = at::kBool; +inline constexpr auto kFloat8_e4m3fn = at::kFloat8_e4m3fn; +inline constexpr auto kFloat8_e5m2 = at::kFloat8_e5m2; + +// --- Factories / ops re-exported with identical semantics ------------------ +// (using-declarations bring every overload; behaviour is unchanged.) +// NOTE: at::device()/at::dtype() are intentionally NOT re-exported -- their bare +// names collide with the many locals/params called `device`/`dtype`. Build +// TensorOptions explicitly instead: TensorOptions().device(...).dtype(...). +using at::CUDA; +using at::empty; +using at::empty_like; +using at::from_blob; +using at::get_generator_or_default; +using at::reciprocal; +using at::sum_out; +using at::zeros; +using c10::elementSize; +using torch::range; + +// --- CUDA helpers ---------------------------------------------------------- +using at::cuda::current_device; +using at::cuda::getCurrentCUDAStream; +using at::cuda::getCurrentDeviceProperties; +using at::cuda::getStreamFromExternal; +using at::cuda::detail::getDefaultCUDAGenerator; + +// --- torch.Tensor indexing (Slice/None/TensorIndex) ------------------------ +namespace indexing = torch::indexing; + +#endif + +// Optional tensor, matches python's ``Optional[torch.Tensor]``. +using MaybeTensor = std::optional; + +// =========================================================================== +// dtype mapping -- the TE<->torch dtype boundary. +// =========================================================================== + +// These are defined inline (in this facade header) on purpose: they sit on the +// per-tensor Python<->C++ marshalling path, so keeping them inlinable at call +// sites avoids an out-of-line call and matches the pre-facade codegen. + +/*! \brief Map a TE DType to the corresponding torch scalar type. */ +inline ScalarType GetATenDType(transformer_engine::DType t) { + switch (t) { + case transformer_engine::DType::kInt16: + return torch::kInt16; + case transformer_engine::DType::kInt32: + return torch::kInt32; + case transformer_engine::DType::kInt64: + return torch::kInt64; + case transformer_engine::DType::kFloat32: + return at::kFloat; + case transformer_engine::DType::kFloat16: + return at::kHalf; + case transformer_engine::DType::kBFloat16: + return at::kBFloat16; + case transformer_engine::DType::kByte: + return at::kByte; + case transformer_engine::DType::kFloat8E4M3: + return at::kFloat8_e4m3fn; + case transformer_engine::DType::kFloat8E5M2: + return at::kFloat8_e5m2; + case transformer_engine::DType::kFloat8E8M0: + return at::kByte; // e8m0 dtype requires PyTorch 2.7.0+ + default: + NVTE_ERROR("Invalid type (", static_cast(t), ")."); + } +} + +/*! \brief Map a torch scalar type to the corresponding TE DType. */ +inline transformer_engine::DType GetTransformerEngineDType(ScalarType t) { + switch (t) { + case at::kFloat8_e4m3fn: + return transformer_engine::DType::kFloat8E4M3; + case at::kFloat8_e5m2: + return transformer_engine::DType::kFloat8E5M2; + case at::kHalf: + return transformer_engine::DType::kFloat16; + case at::kFloat: + return transformer_engine::DType::kFloat32; + case at::kBFloat16: + return transformer_engine::DType::kBFloat16; + case at::kBool: + return transformer_engine::DType::kByte; + case torch::kByte: + return transformer_engine::DType::kByte; + case torch::kInt16: + return transformer_engine::DType::kInt16; + case torch::kInt32: + return transformer_engine::DType::kInt32; + case torch::kInt64: + return transformer_engine::DType::kInt64; + default: + NVTE_ERROR("Invalid type (", static_cast(t), ")."); + } +} + +// =========================================================================== +// Tensor factories -- wrappers over at::empty/at::zeros/... so no extension +// file has to name the ATen factory functions. +// =========================================================================== + +/*! \brief Allocate a CUDA tensor of the given shape/dtype. + * \param zero_init If true, zero-initialize; otherwise leave uninitialized. + */ +inline Tensor new_cuda_tensor(const std::vector& shape, ScalarType dtype, bool zero_init) { + c10::IntArrayRef ar_shape(shape); + if (zero_init) { + return at::zeros(ar_shape, at::CUDA(dtype)); + } + return at::empty(ar_shape, at::CUDA(dtype)); +} + +} // namespace transformer_engine::pytorch + +#endif // TRANSFORMER_ENGINE_PYTORCH_CSRC_TORCH_BACKEND_H_ diff --git a/transformer_engine/pytorch/csrc/type_converters.cpp b/transformer_engine/pytorch/csrc/type_converters.cpp index ddb85808a5..3771c53fd5 100644 --- a/transformer_engine/pytorch/csrc/type_converters.cpp +++ b/transformer_engine/pytorch/csrc/type_converters.cpp @@ -4,7 +4,6 @@ * See LICENSE for license information. ************************************************************************/ -#include #include #include @@ -26,19 +25,19 @@ TensorWrapper NVTETensorFromFloat8Tensor(py::handle tensor, Quantizer *quantizer // FP8 data const DType fp8_dtype = tensor.attr("_fp8_dtype").cast(); if (data_exists) { - const auto &data = tensor.attr("_data").cast(); + const auto &data = tensor.attr("_data").cast(); ret.set_rowwise_data(data.data_ptr(), fp8_dtype, getTensorShape(data)); } // FP8 data transpose if (transpose_exists) { - const auto &data_transpose = tensor.attr("_transpose").cast(); + const auto &data_transpose = tensor.attr("_transpose").cast(); ret.set_columnwise_data(data_transpose.data_ptr(), fp8_dtype, getTensorShape(data_transpose)); } // Scale-inverse { - const auto &scale_inv = tensor.attr("_scale_inv").cast(); + const auto &scale_inv = tensor.attr("_scale_inv").cast(); float *dptr = reinterpret_cast(scale_inv.data_ptr()); const auto &dtype = GetTransformerEngineDType(scale_inv.scalar_type()); const auto &shape = getTensorShape(scale_inv); @@ -64,16 +63,16 @@ TensorWrapper NVTETensorFromMXFP8Tensor(py::handle tensor, Quantizer *quantizer) // Row-scaled data const DType fp8_dtype = tensor.attr("_fp8_dtype").cast(); if (rowwise_usage) { - const auto &data = tensor.attr("_rowwise_data").cast(); - const auto &scale_inv = tensor.attr("_rowwise_scale_inv").cast(); + const auto &data = tensor.attr("_rowwise_data").cast(); + const auto &scale_inv = tensor.attr("_rowwise_scale_inv").cast(); ret.set_rowwise_data(data.data_ptr(), fp8_dtype, getTensorShape(data)); ret.set_rowwise_scale_inv(scale_inv.data_ptr(), DType::kFloat8E8M0, getTensorShape(scale_inv)); } // Column-scaled data if (columnwise_usage) { - const auto &data = tensor.attr("_columnwise_data").cast(); - const auto &scale_inv = tensor.attr("_columnwise_scale_inv").cast(); + const auto &data = tensor.attr("_columnwise_data").cast(); + const auto &scale_inv = tensor.attr("_columnwise_scale_inv").cast(); ret.set_columnwise_data(data.data_ptr(), fp8_dtype, getTensorShape(data)); ret.set_columnwise_scale_inv(scale_inv.data_ptr(), DType::kFloat8E8M0, getTensorShape(scale_inv)); @@ -99,8 +98,8 @@ TensorWrapper NVTETensorFromFloat8BlockwiseQTensor(py::handle tensor, Quantizer // Row-wise data if (rowwise_usage) { - const at::Tensor &data_rowwise = tensor.attr("_rowwise_data").cast(); - const at::Tensor &scale_inv_rowwise = tensor.attr("_rowwise_scale_inv").cast(); + const Tensor &data_rowwise = tensor.attr("_rowwise_data").cast(); + const Tensor &scale_inv_rowwise = tensor.attr("_rowwise_scale_inv").cast(); void *scale_inv_rowwise_dptr = scale_inv_rowwise.data_ptr(); const auto &rowwise_shape = getTensorShape(data_rowwise); ret.set_rowwise_data(data_rowwise.data_ptr(), dtype, rowwise_shape); @@ -110,8 +109,8 @@ TensorWrapper NVTETensorFromFloat8BlockwiseQTensor(py::handle tensor, Quantizer // Column-wise data if (columnwise_usage) { - const at::Tensor &data_colwise = tensor.attr("_columnwise_data").cast(); - const at::Tensor &scale_inv_colwise = tensor.attr("_columnwise_scale_inv").cast(); + const Tensor &data_colwise = tensor.attr("_columnwise_data").cast(); + const Tensor &scale_inv_colwise = tensor.attr("_columnwise_scale_inv").cast(); void *scale_inv_colwise_dptr = scale_inv_colwise.data_ptr(); const auto &shape = getTensorShape(data_colwise); ret.set_columnwise_data(data_colwise.data_ptr(), dtype, shape); @@ -141,9 +140,9 @@ TensorWrapper NVTETensorFromNVFP4Tensor(py::handle tensor, Quantizer *quantizer) // Row-scaled data if (rowwise_usage) { - const auto &data = tensor.attr("_rowwise_data").cast(); - const auto &scale_inv = tensor.attr("_rowwise_scale_inv").cast(); - const auto &amax_rowwise = tensor.attr("_amax_rowwise").cast(); + const auto &data = tensor.attr("_rowwise_data").cast(); + const auto &scale_inv = tensor.attr("_rowwise_scale_inv").cast(); + const auto &amax_rowwise = tensor.attr("_amax_rowwise").cast(); ret.set_rowwise_data(data.data_ptr(), dtype, convert_shape_back_from_fp4(getTensorShape(data), false)); ret.set_rowwise_scale_inv(scale_inv.data_ptr(), DType::kFloat8E4M3, getTensorShape(scale_inv)); @@ -152,9 +151,9 @@ TensorWrapper NVTETensorFromNVFP4Tensor(py::handle tensor, Quantizer *quantizer) // Column-scaled data if (columnwise_usage) { - const auto &data = tensor.attr("_columnwise_data").cast(); - const auto &scale_inv = tensor.attr("_columnwise_scale_inv").cast(); - const auto &amax_columnwise = tensor.attr("_amax_columnwise").cast(); + const auto &data = tensor.attr("_columnwise_data").cast(); + const auto &scale_inv = tensor.attr("_columnwise_scale_inv").cast(); + const auto &amax_columnwise = tensor.attr("_amax_columnwise").cast(); ret.set_columnwise_data(data.data_ptr(), DType::kFloat4E2M1, convert_shape_back_from_fp4(getTensorShape(data), false)); ret.set_columnwise_scale_inv(scale_inv.data_ptr(), DType::kFloat8E4M3, @@ -189,7 +188,7 @@ NVTEScalingMode ScalingModeFromQuantizer(py::handle quantizer) { return NVTE_DELAYED_TENSOR_SCALING; } -DType GetTransformerEngineDTypeForScaleInv(py::handle quantizer, at::Tensor scale_inv) { +DType GetTransformerEngineDTypeForScaleInv(py::handle quantizer, Tensor scale_inv) { auto *quantizer_ptr = quantizer.ptr(); if (IsMXFP8Quantizers(quantizer_ptr)) { return DType::kFloat8E8M0; @@ -221,7 +220,7 @@ GroupedTensorWrapper GroupedTensorFromPyTorchGroupedTensor(py::handle tensor) { // Rowwise data if (!tensor.attr("rowwise_data").is_none()) { - const auto &data = tensor.attr("rowwise_data").cast(); + const auto &data = tensor.attr("rowwise_data").cast(); DType data_dtype = quantizer.is_none() ? GetTransformerEngineDType(data.scalar_type()) : quantizer_dtype; ret.set_rowwise_data(data.data_ptr(), data_dtype, getTensorShape(data)); @@ -231,7 +230,7 @@ GroupedTensorWrapper GroupedTensorFromPyTorchGroupedTensor(py::handle tensor) { // Columnwise data if (!tensor.attr("columnwise_data").is_none()) { - const auto &data = tensor.attr("columnwise_data").cast(); + const auto &data = tensor.attr("columnwise_data").cast(); DType data_dtype = quantizer.is_none() ? GetTransformerEngineDType(data.scalar_type()) : quantizer_dtype; ret.set_columnwise_data(data.data_ptr(), data_dtype, getTensorShape(data)); @@ -241,32 +240,32 @@ GroupedTensorWrapper GroupedTensorFromPyTorchGroupedTensor(py::handle tensor) { // Scale if (!tensor.attr("scale").is_none()) { - const auto &scale = tensor.attr("scale").cast(); + const auto &scale = tensor.attr("scale").cast(); ret.set_scale(scale.data_ptr(), GetTransformerEngineDType(scale.scalar_type()), getTensorShape(scale)); } // Amax if (!tensor.attr("amax").is_none()) { - const auto &amax = tensor.attr("amax").cast(); + const auto &amax = tensor.attr("amax").cast(); ret.set_amax(amax.data_ptr(), GetTransformerEngineDType(amax.scalar_type()), getTensorShape(amax)); } if (!tensor.attr("columnwise_amax").is_none()) { - const auto &amax = tensor.attr("columnwise_amax").cast(); + const auto &amax = tensor.attr("columnwise_amax").cast(); ret.set_columnwise_amax(amax.data_ptr(), GetTransformerEngineDType(amax.scalar_type()), getTensorShape(amax)); } // Scale inverse if (!tensor.attr("scale_inv").is_none()) { - const auto &scale_inv = tensor.attr("scale_inv").cast(); + const auto &scale_inv = tensor.attr("scale_inv").cast(); ret.set_rowwise_scale_inv(scale_inv.data_ptr(), GetTransformerEngineDTypeForScaleInv(quantizer, scale_inv), getTensorShape(scale_inv)); } if (!tensor.attr("columnwise_scale_inv").is_none()) { - const auto &scale_inv = tensor.attr("columnwise_scale_inv").cast(); + const auto &scale_inv = tensor.attr("columnwise_scale_inv").cast(); ret.set_columnwise_scale_inv(scale_inv.data_ptr(), GetTransformerEngineDTypeForScaleInv(quantizer, scale_inv), getTensorShape(scale_inv)); @@ -274,17 +273,17 @@ GroupedTensorWrapper GroupedTensorFromPyTorchGroupedTensor(py::handle tensor) { // Shape metadata if (!tensor.attr("first_dims").is_none()) { - const auto &first_dims = tensor.attr("first_dims").cast(); + const auto &first_dims = tensor.attr("first_dims").cast(); ret.set_first_dims(first_dims.data_ptr(), GetTransformerEngineDType(first_dims.scalar_type()), getTensorShape(first_dims)); } if (!tensor.attr("last_dims").is_none()) { - const auto &last_dims = tensor.attr("last_dims").cast(); + const auto &last_dims = tensor.attr("last_dims").cast(); ret.set_last_dims(last_dims.data_ptr(), GetTransformerEngineDType(last_dims.scalar_type()), getTensorShape(last_dims)); } if (!tensor.attr("tensor_offsets").is_none()) { - const auto &tensor_offsets = tensor.attr("tensor_offsets").cast(); + const auto &tensor_offsets = tensor.attr("tensor_offsets").cast(); ret.set_tensor_offsets(tensor_offsets.data_ptr(), GetTransformerEngineDType(tensor_offsets.scalar_type()), getTensorShape(tensor_offsets)); diff --git a/transformer_engine/pytorch/csrc/util.h b/transformer_engine/pytorch/csrc/util.h index 132db4075f..e3c6cf1ef5 100644 --- a/transformer_engine/pytorch/csrc/util.h +++ b/transformer_engine/pytorch/csrc/util.h @@ -7,12 +7,11 @@ #ifndef TRANSFORMER_ENGINE_PYTORCH_CSRC_UTIL_H_ #define TRANSFORMER_ENGINE_PYTORCH_CSRC_UTIL_H_ -#include - #include #include #include +#include "common.h" #include "transformer_engine/transformer_engine.h" namespace transformer_engine { @@ -22,21 +21,21 @@ namespace pytorch { * * The returned swizzled scales should be kept alive during the GEMM. */ -std::tuple, std::optional> swizzle_scales_for_gemm( +std::tuple, std::optional> swizzle_scales_for_gemm( TensorWrapper& tensor, bool rowwise_usage, bool columnwise_usage); /*! \brief Convert multiple tensor block scales into GEMM swizzled format. * * The returned swizzled scales should be kept alive during the GEMMs. */ -std::optional multi_tensor_swizzle_scales_for_gemm(std::vector& tensors, +std::optional multi_tensor_swizzle_scales_for_gemm(std::vector& tensors, bool rowwise_usage, bool columnwise_usage); -std::optional multi_tensor_swizzle_scales_for_gemm_unchecked( +std::optional multi_tensor_swizzle_scales_for_gemm_unchecked( std::vector& tensors, bool rowwise_usage, bool columnwise_usage); -using SwizzledGroupedScales = std::pair, std::optional>; +using SwizzledGroupedScales = std::pair, std::optional>; /*! \brief Swizzle grouped tensor scales for GEMM if needed. * Currently only works for MXFP8 1D scaling with uniform shapes. @@ -63,7 +62,7 @@ std::optional maybe_swizzle_grouped_tensor(GroupedTensorW * The returned swizzled scaling factor tensor should be kept alive * during the GEMM. */ -at::Tensor convert_block_scaling_to_mxfp8_tensor(TensorWrapper& input, bool rowwise); +Tensor convert_block_scaling_to_mxfp8_tensor(TensorWrapper& input, bool rowwise); } // namespace pytorch } // namespace transformer_engine