From 05df7232c7864fe870353343e9d9a270e1e36a1e Mon Sep 17 00:00:00 2001 From: Jiacheng Huang Date: Wed, 5 Aug 2026 10:53:35 +0800 Subject: [PATCH] fix: copy optional tensor metadata --- src/base/binary_cross_entropy.h | 6 ++++-- src/base/flash_attn_varlen_func.h | 17 ++++++++++------- src/base/flash_attn_with_kvcache.h | 10 ++++++---- src/base/multi_margin_loss.h | 6 ++++-- 4 files changed, 24 insertions(+), 15 deletions(-) diff --git a/src/base/binary_cross_entropy.h b/src/base/binary_cross_entropy.h index 3f9f10d49..5d743a9d8 100644 --- a/src/base/binary_cross_entropy.h +++ b/src/base/binary_cross_entropy.h @@ -26,8 +26,10 @@ class BinaryCrossEntropy : public Operator { out_strides_{out.strides()}, out_type_{out.dtype()}, has_weight_{weight.has_value()}, - weight_shape_{weight ? weight->shape() : Tensor::Shape{}}, - weight_strides_{weight ? weight->strides() : Tensor::Strides{}}, + weight_shape_{weight ? Tensor::Shape{weight->shape()} + : Tensor::Shape{}}, + weight_strides_{weight ? Tensor::Strides{weight->strides()} + : Tensor::Strides{}}, weight_type_{weight ? weight->dtype() : DataType::kFloat32}, reduction_{reduction_detail::FromPythonArguments(size_average, reduce, reduction)}, diff --git a/src/base/flash_attn_varlen_func.h b/src/base/flash_attn_varlen_func.h index c645302e5..dda248450 100644 --- a/src/base/flash_attn_varlen_func.h +++ b/src/base/flash_attn_varlen_func.h @@ -53,9 +53,10 @@ class FlashAttnVarlenFunc : public Operator { cu_seqlens_q_shape_{cu_seqlens_q.shape()}, cu_seqlens_k_shape_{cu_seqlens_k.shape()}, out_shape_{out.shape()}, - softmax_lse_shape_{softmax_lse.has_value() ? softmax_lse->shape() - : Tensor::Shape{}}, - s_dmask_shape_{s_dmask.has_value() ? s_dmask->shape() + softmax_lse_shape_{softmax_lse.has_value() + ? Tensor::Shape{softmax_lse->shape()} + : Tensor::Shape{}}, + s_dmask_shape_{s_dmask.has_value() ? Tensor::Shape{s_dmask->shape()} : Tensor::Shape{}}, q_strides_{q.strides()}, k_strides_{k.strides()}, @@ -63,10 +64,12 @@ class FlashAttnVarlenFunc : public Operator { cu_seqlens_q_strides_{cu_seqlens_q.strides()}, cu_seqlens_k_strides_{cu_seqlens_k.strides()}, out_strides_{out.strides()}, - softmax_lse_strides_{softmax_lse.has_value() ? softmax_lse->strides() - : Tensor::Strides{}}, - s_dmask_strides_{s_dmask.has_value() ? s_dmask->strides() - : Tensor::Strides{}}, + softmax_lse_strides_{softmax_lse.has_value() + ? Tensor::Strides{softmax_lse->strides()} + : Tensor::Strides{}}, + s_dmask_strides_{s_dmask.has_value() + ? Tensor::Strides{s_dmask->strides()} + : Tensor::Strides{}}, q_dtype_{q.dtype()}, k_dtype_{k.dtype()}, v_dtype_{v.dtype()}, diff --git a/src/base/flash_attn_with_kvcache.h b/src/base/flash_attn_with_kvcache.h index dbc37207a..ecd2c14c9 100644 --- a/src/base/flash_attn_with_kvcache.h +++ b/src/base/flash_attn_with_kvcache.h @@ -123,8 +123,9 @@ class FlashAttnWithKvcache : public Operator { ? Tensor::Shape{alibi_slopes->shape()} : Tensor::Shape{}}, out_shape_{out.shape()}, - softmax_lse_shape_{softmax_lse.has_value() ? softmax_lse->shape() - : Tensor::Shape{}}, + softmax_lse_shape_{softmax_lse.has_value() + ? Tensor::Shape{softmax_lse->shape()} + : Tensor::Shape{}}, q_strides_{q.strides()}, k_cache_strides_{k_cache.strides()}, v_cache_strides_{v_cache.strides()}, @@ -155,8 +156,9 @@ class FlashAttnWithKvcache : public Operator { ? Tensor::Strides{alibi_slopes->strides()} : Tensor::Strides{}}, out_strides_{out.strides()}, - softmax_lse_strides_{softmax_lse.has_value() ? softmax_lse->strides() - : Tensor::Strides{}}, + softmax_lse_strides_{softmax_lse.has_value() + ? Tensor::Strides{softmax_lse->strides()} + : Tensor::Strides{}}, q_dtype_{q.dtype()}, k_cache_dtype_{k_cache.dtype()}, v_cache_dtype_{v_cache.dtype()}, diff --git a/src/base/multi_margin_loss.h b/src/base/multi_margin_loss.h index 7a024bf00..503cc1e00 100644 --- a/src/base/multi_margin_loss.h +++ b/src/base/multi_margin_loss.h @@ -26,8 +26,10 @@ class MultiMarginLoss : public Operator { out_strides_{out.strides()}, out_type_{out.dtype()}, has_weight_{weight.has_value()}, - weight_shape_{weight ? weight->shape() : Tensor::Shape{}}, - weight_strides_{weight ? weight->strides() : Tensor::Strides{}}, + weight_shape_{weight ? Tensor::Shape{weight->shape()} + : Tensor::Shape{}}, + weight_strides_{weight ? Tensor::Strides{weight->strides()} + : Tensor::Strides{}}, weight_type_{weight ? weight->dtype() : DataType::kFloat32}, p_{static_cast(p)}, margin_{margin},