Kernels:
Trusted publisher
Uploaded using `kernel-builder`.
Browse files- build/torch-rocm/_ops.py +1 -1
- build/torch-rocm/cross_entropy.py +54 -50
- build/torch-rocm/dyt.py +16 -12
- build/torch-rocm/fused_linear_cross_entropy.py +72 -69
- build/torch-rocm/geglu.py +22 -18
- build/torch-rocm/group_norm.py +43 -39
- build/torch-rocm/jsd.py +19 -17
- build/torch-rocm/kl_div.py +28 -25
- build/torch-rocm/layer_norm.py +47 -43
- build/torch-rocm/metadata.json +17 -17
- build/torch-rocm/metadata.json.sigstore +1 -1
- build/torch-rocm/qwen2vl_mrope.py +41 -37
- build/torch-rocm/rms_norm.py +97 -92
- build/torch-rocm/rope.py +47 -43
- build/torch-rocm/swiglu.py +23 -20
- build/torch-rocm/tvd.py +20 -18
- build/torch-rocm/utils.py +11 -0
build/torch-rocm/_ops.py
CHANGED
|
@@ -22,7 +22,7 @@ def get_backend() -> str:
|
|
| 22 |
|
| 23 |
def _find_ops_name() -> str:
|
| 24 |
kernel_name = "liger_kernels"
|
| 25 |
-
unique_id = "
|
| 26 |
backend = get_backend()
|
| 27 |
return f"_{kernel_name}_{backend}_{unique_id}"
|
| 28 |
|
|
|
|
| 22 |
|
| 23 |
def _find_ops_name() -> str:
|
| 24 |
kernel_name = "liger_kernels"
|
| 25 |
+
unique_id = "0c5fb33"
|
| 26 |
backend = get_backend()
|
| 27 |
return f"_{kernel_name}_{backend}_{unique_id}"
|
| 28 |
|
build/torch-rocm/cross_entropy.py
CHANGED
|
@@ -11,6 +11,8 @@ from .utils import element_mul_kernel
|
|
| 11 |
from .utils import is_hip
|
| 12 |
from .utils import infer_device
|
| 13 |
from .utils import is_npu_available
|
|
|
|
|
|
|
| 14 |
|
| 15 |
if compare_version("triton", operator.ge, "3.0.0") and not is_npu_available():
|
| 16 |
try:
|
|
@@ -372,44 +374,45 @@ def cross_entropy_forward(
|
|
| 372 |
if target.stride(-1) != 1:
|
| 373 |
target = target.contiguous()
|
| 374 |
|
| 375 |
-
|
| 376 |
-
|
| 377 |
-
|
| 378 |
-
|
| 379 |
-
|
| 380 |
-
|
| 381 |
-
|
| 382 |
-
|
| 383 |
-
|
| 384 |
-
|
| 385 |
-
|
| 386 |
-
|
| 387 |
-
|
| 388 |
-
|
| 389 |
-
|
| 390 |
-
|
| 391 |
-
|
| 392 |
-
|
| 393 |
-
|
| 394 |
-
|
| 395 |
-
|
| 396 |
-
|
| 397 |
-
|
| 398 |
-
|
| 399 |
-
|
| 400 |
-
|
| 401 |
-
|
| 402 |
-
|
| 403 |
-
|
| 404 |
-
|
| 405 |
-
|
| 406 |
-
|
| 407 |
-
|
| 408 |
-
|
| 409 |
-
|
| 410 |
-
|
| 411 |
-
|
| 412 |
-
|
|
|
|
| 413 |
|
| 414 |
if reduction == "none":
|
| 415 |
loss = loss_1d
|
|
@@ -437,18 +440,19 @@ def cross_entropy_backward(_input, grad_output):
|
|
| 437 |
# We use a Triton kernel instead of a PyTorch operation because modifying inputs in-place
|
| 438 |
# for gradient storage and backward multiple times causes anomalies with PyTorch but not with Triton.
|
| 439 |
else:
|
| 440 |
-
|
| 441 |
-
|
| 442 |
-
|
| 443 |
-
|
| 444 |
-
|
| 445 |
-
|
| 446 |
-
|
| 447 |
-
|
| 448 |
-
|
| 449 |
-
|
| 450 |
-
|
| 451 |
-
|
|
|
|
| 452 |
|
| 453 |
return _input
|
| 454 |
|
|
|
|
| 11 |
from .utils import is_hip
|
| 12 |
from .utils import infer_device
|
| 13 |
from .utils import is_npu_available
|
| 14 |
+
from .utils import device_context
|
| 15 |
+
|
| 16 |
|
| 17 |
if compare_version("triton", operator.ge, "3.0.0") and not is_npu_available():
|
| 18 |
try:
|
|
|
|
| 374 |
if target.stride(-1) != 1:
|
| 375 |
target = target.contiguous()
|
| 376 |
|
| 377 |
+
with device_context(_input.device):
|
| 378 |
+
# Here we use a trick to store X_ptr gradient in X_ptr so we can save memory
|
| 379 |
+
liger_cross_entropy_kernel[(n_rows,)](
|
| 380 |
+
X_ptr=_input,
|
| 381 |
+
X_stride=_input.stride(-2),
|
| 382 |
+
Y_ptr=target,
|
| 383 |
+
Y_stride=target.stride(-1), # always 1
|
| 384 |
+
weight_ptr=weight, # dummy if None
|
| 385 |
+
loss_ptr=loss_1d,
|
| 386 |
+
z_loss_ptr=z_loss_1d,
|
| 387 |
+
loss_stride=loss_1d.stride(-1), # always 1
|
| 388 |
+
token_accuracy_ptr=token_accuracy_1d,
|
| 389 |
+
token_accuracy_stride=token_accuracy_1d.stride(-1)
|
| 390 |
+
if return_token_accuracy
|
| 391 |
+
else 0, # always 1 if accuracy is enabled
|
| 392 |
+
predicted_tokens_ptr=predicted_tokens_1d,
|
| 393 |
+
predicted_tokens_stride=predicted_tokens_1d.stride(-1)
|
| 394 |
+
if return_predicted_tokens
|
| 395 |
+
else 0, # always 1 if predicted tokens is enabled
|
| 396 |
+
n_cols=V,
|
| 397 |
+
n_non_ignore=n_non_ignore,
|
| 398 |
+
sum_non_ignore_weight=sum_non_ignore_weight,
|
| 399 |
+
ignore_index=ignore_index,
|
| 400 |
+
weight_sum=weight_sum,
|
| 401 |
+
lse_square_scale=lse_square_scale,
|
| 402 |
+
label_smoothing=label_smoothing,
|
| 403 |
+
reduction=reduction,
|
| 404 |
+
softcap=softcap,
|
| 405 |
+
RETURN_Z_LOSS=return_z_loss,
|
| 406 |
+
RETURN_TOKEN_ACCURACY=return_token_accuracy,
|
| 407 |
+
RETURN_PREDICTED_TOKENS=return_predicted_tokens,
|
| 408 |
+
BLOCK_SIZE=BLOCK_SIZE,
|
| 409 |
+
HAS_WEIGHT=True if weight is not None else False,
|
| 410 |
+
HAS_SOFTCAPPING=True if softcap is not None else False,
|
| 411 |
+
HAS_GRADIENTS=_input.requires_grad,
|
| 412 |
+
# TODO: 32 seems to give the best performance
|
| 413 |
+
# Performance is quite sensitive to num_warps
|
| 414 |
+
num_warps=32 if not is_hip() else 16,
|
| 415 |
+
)
|
| 416 |
|
| 417 |
if reduction == "none":
|
| 418 |
loss = loss_1d
|
|
|
|
| 440 |
# We use a Triton kernel instead of a PyTorch operation because modifying inputs in-place
|
| 441 |
# for gradient storage and backward multiple times causes anomalies with PyTorch but not with Triton.
|
| 442 |
else:
|
| 443 |
+
with device_context(_input.device):
|
| 444 |
+
BT, V = _input.shape
|
| 445 |
+
n_rows = BT
|
| 446 |
+
BLOCK_SIZE = min(MAX_FUSED_SIZE, triton.next_power_of_2(V))
|
| 447 |
+
|
| 448 |
+
element_mul_kernel[(n_rows,)](
|
| 449 |
+
_input,
|
| 450 |
+
_input.stride(-2),
|
| 451 |
+
grad_output,
|
| 452 |
+
V,
|
| 453 |
+
BLOCK_SIZE=BLOCK_SIZE,
|
| 454 |
+
num_warps=32 if not is_hip() else 16,
|
| 455 |
+
)
|
| 456 |
|
| 457 |
return _input
|
| 458 |
|
build/torch-rocm/dyt.py
CHANGED
|
@@ -9,6 +9,8 @@ from .utils import ensure_contiguous
|
|
| 9 |
from .utils import get_npu_core_count
|
| 10 |
from .utils import infer_device
|
| 11 |
from .utils import is_npu_available
|
|
|
|
|
|
|
| 12 |
|
| 13 |
if compare_version("triton", operator.ge, "3.0.0") and not is_npu_available():
|
| 14 |
try:
|
|
@@ -107,16 +109,17 @@ def liger_dyt_fwd(x, alpha, gamma, beta):
|
|
| 107 |
|
| 108 |
y = torch.empty_like(x)
|
| 109 |
|
| 110 |
-
|
| 111 |
-
|
| 112 |
-
|
| 113 |
-
|
| 114 |
-
|
| 115 |
-
|
| 116 |
-
|
| 117 |
-
|
| 118 |
-
|
| 119 |
-
|
|
|
|
| 120 |
return y.view(input_shape)
|
| 121 |
|
| 122 |
|
|
@@ -139,8 +142,9 @@ def liger_dyt_bwd(dy, x, alpha, gamma, beta):
|
|
| 139 |
db = torch.empty(NUM_SMS, N, dtype=torch.float32, device=x.device) if HAVE_BETA else None
|
| 140 |
dx = torch.empty_like(dy)
|
| 141 |
|
| 142 |
-
|
| 143 |
-
|
|
|
|
| 144 |
if HAVE_BETA:
|
| 145 |
db = db.sum(0).to(x.dtype)
|
| 146 |
dg = dg.sum(0).to(gamma.dtype)
|
|
|
|
| 9 |
from .utils import get_npu_core_count
|
| 10 |
from .utils import infer_device
|
| 11 |
from .utils import is_npu_available
|
| 12 |
+
from .utils import device_context
|
| 13 |
+
|
| 14 |
|
| 15 |
if compare_version("triton", operator.ge, "3.0.0") and not is_npu_available():
|
| 16 |
try:
|
|
|
|
| 109 |
|
| 110 |
y = torch.empty_like(x)
|
| 111 |
|
| 112 |
+
with device_context(x.device):
|
| 113 |
+
grid = lambda meta: (triton.cdiv(N, meta["BLOCK_N"]), M)
|
| 114 |
+
_dyt_fwd_kernel[grid](
|
| 115 |
+
x,
|
| 116 |
+
y,
|
| 117 |
+
alpha,
|
| 118 |
+
gamma,
|
| 119 |
+
beta,
|
| 120 |
+
HAVE_BETA,
|
| 121 |
+
N,
|
| 122 |
+
)
|
| 123 |
return y.view(input_shape)
|
| 124 |
|
| 125 |
|
|
|
|
| 142 |
db = torch.empty(NUM_SMS, N, dtype=torch.float32, device=x.device) if HAVE_BETA else None
|
| 143 |
dx = torch.empty_like(dy)
|
| 144 |
|
| 145 |
+
with device_context(x.device):
|
| 146 |
+
grid = lambda meta: (triton.cdiv(N, meta["BLOCK_N"]), NUM_SMS)
|
| 147 |
+
_dyt_bwd_kernel[grid](dy, dx, da, dg, db, x, alpha, gamma, HAVE_BETA, M, N)
|
| 148 |
if HAVE_BETA:
|
| 149 |
db = db.sum(0).to(x.dtype)
|
| 150 |
dg = dg.sum(0).to(gamma.dtype)
|
build/torch-rocm/fused_linear_cross_entropy.py
CHANGED
|
@@ -7,6 +7,7 @@ from .utils import amp_custom_fwd
|
|
| 7 |
from .utils import element_mul_kernel
|
| 8 |
from .utils import is_hip
|
| 9 |
from .utils import infer_device
|
|
|
|
| 10 |
|
| 11 |
# The hard limit of TRITON_MAX_TENSOR_NUMEL is 1048576 https://github.com/triton-lang/triton/blob/ba42a5c68fd0505f8c42f4202d53be0f8d9a5fe0/python/triton/language/core.py#L19
|
| 12 |
# However, setting limit as 65536 as in LayerNorm tutorial is faster because of less register spilling
|
|
@@ -147,42 +148,43 @@ def fused_linear_cross_entropy_forward(
|
|
| 147 |
logits_chunk = logits_chunk.contiguous()
|
| 148 |
target_chunk = target_chunk.contiguous()
|
| 149 |
|
| 150 |
-
|
| 151 |
-
|
| 152 |
-
|
| 153 |
-
|
| 154 |
-
|
| 155 |
-
|
| 156 |
-
|
| 157 |
-
|
| 158 |
-
|
| 159 |
-
|
| 160 |
-
|
| 161 |
-
|
| 162 |
-
|
| 163 |
-
|
| 164 |
-
|
| 165 |
-
|
| 166 |
-
|
| 167 |
-
|
| 168 |
-
|
| 169 |
-
|
| 170 |
-
|
| 171 |
-
|
| 172 |
-
|
| 173 |
-
|
| 174 |
-
|
| 175 |
-
|
| 176 |
-
|
| 177 |
-
|
| 178 |
-
|
| 179 |
-
|
| 180 |
-
|
| 181 |
-
|
| 182 |
-
|
| 183 |
-
|
| 184 |
-
|
| 185 |
-
|
|
|
|
| 186 |
|
| 187 |
# Apply token scaling if requested
|
| 188 |
if use_token_scaling:
|
|
@@ -247,47 +249,48 @@ def fused_linear_cross_entropy_forward(
|
|
| 247 |
def fused_linear_cross_entropy_backward(grad_output, grad_input, grad_weight, grad_bias):
|
| 248 |
# If cross entropy is the last layer, grad_output is 1.0. Skip the mul to save time
|
| 249 |
if not torch.equal(grad_output, torch.tensor(1.0, device=grad_output.device)):
|
| 250 |
-
|
| 251 |
-
|
| 252 |
-
|
| 253 |
-
|
| 254 |
-
|
| 255 |
-
|
| 256 |
-
element_mul_kernel[(n_rows,)](
|
| 257 |
-
grad_input,
|
| 258 |
-
grad_input.stride(-2),
|
| 259 |
-
grad_output,
|
| 260 |
-
H,
|
| 261 |
-
BLOCK_SIZE=BLOCK_SIZE,
|
| 262 |
-
num_warps=32 if not is_hip() else 16,
|
| 263 |
-
)
|
| 264 |
-
|
| 265 |
-
# handle grad_weight
|
| 266 |
-
if grad_weight is not None:
|
| 267 |
-
V, H = grad_weight.shape
|
| 268 |
-
n_rows = V
|
| 269 |
|
| 270 |
element_mul_kernel[(n_rows,)](
|
| 271 |
-
|
| 272 |
-
|
| 273 |
grad_output,
|
| 274 |
H,
|
| 275 |
BLOCK_SIZE=BLOCK_SIZE,
|
| 276 |
num_warps=32 if not is_hip() else 16,
|
| 277 |
)
|
| 278 |
|
| 279 |
-
|
| 280 |
-
|
| 281 |
-
|
| 282 |
-
|
| 283 |
-
|
| 284 |
-
|
| 285 |
-
|
| 286 |
-
|
| 287 |
-
|
| 288 |
-
|
| 289 |
-
|
| 290 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 291 |
return grad_input, grad_weight, grad_bias
|
| 292 |
|
| 293 |
|
|
|
|
| 7 |
from .utils import element_mul_kernel
|
| 8 |
from .utils import is_hip
|
| 9 |
from .utils import infer_device
|
| 10 |
+
from .utils import device_context
|
| 11 |
|
| 12 |
# The hard limit of TRITON_MAX_TENSOR_NUMEL is 1048576 https://github.com/triton-lang/triton/blob/ba42a5c68fd0505f8c42f4202d53be0f8d9a5fe0/python/triton/language/core.py#L19
|
| 13 |
# However, setting limit as 65536 as in LayerNorm tutorial is faster because of less register spilling
|
|
|
|
| 148 |
logits_chunk = logits_chunk.contiguous()
|
| 149 |
target_chunk = target_chunk.contiguous()
|
| 150 |
|
| 151 |
+
with device_context(device):
|
| 152 |
+
# Here we calculate the gradient of logits_chunk in place so we can save memory.
|
| 153 |
+
liger_cross_entropy_kernel[(n_rows,)](
|
| 154 |
+
X_ptr=logits_chunk,
|
| 155 |
+
X_stride=logits_chunk.stride(-2),
|
| 156 |
+
Y_ptr=target_chunk,
|
| 157 |
+
Y_stride=target_chunk.stride(-1), # always 1
|
| 158 |
+
weight_ptr=ce_weight,
|
| 159 |
+
loss_ptr=loss_1d_slice,
|
| 160 |
+
z_loss_ptr=z_loss_1d_slice,
|
| 161 |
+
loss_stride=loss_1d_slice.stride(-1), # always 1
|
| 162 |
+
token_accuracy_ptr=token_accuracy_1d_slice,
|
| 163 |
+
token_accuracy_stride=token_accuracy_1d_slice.stride(-1)
|
| 164 |
+
if return_token_accuracy
|
| 165 |
+
else 0, # always 1 if accuracy is enabled
|
| 166 |
+
predicted_tokens_ptr=predicted_tokens_1d_slice,
|
| 167 |
+
predicted_tokens_stride=predicted_tokens_1d_slice.stride(-1)
|
| 168 |
+
if return_predicted_tokens
|
| 169 |
+
else 0, # always 1 if predicted tokens is enabled
|
| 170 |
+
n_cols=V,
|
| 171 |
+
n_non_ignore=total_n_non_ignore,
|
| 172 |
+
sum_non_ignore_weight=total_sum_non_ignore_ce_weight,
|
| 173 |
+
weight_sum=ce_weight_sum,
|
| 174 |
+
ignore_index=ignore_index,
|
| 175 |
+
lse_square_scale=lse_square_scale,
|
| 176 |
+
label_smoothing=label_smoothing,
|
| 177 |
+
reduction=reduction,
|
| 178 |
+
softcap=softcap,
|
| 179 |
+
RETURN_Z_LOSS=return_z_loss,
|
| 180 |
+
RETURN_TOKEN_ACCURACY=return_token_accuracy,
|
| 181 |
+
RETURN_PREDICTED_TOKENS=return_predicted_tokens,
|
| 182 |
+
HAS_WEIGHT=True if ce_weight is not None else False,
|
| 183 |
+
HAS_SOFTCAPPING=True if softcap is not None else False,
|
| 184 |
+
HAS_GRADIENTS=input_requires_grad,
|
| 185 |
+
BLOCK_SIZE=BLOCK_SIZE,
|
| 186 |
+
num_warps=32 if not is_hip() else 16,
|
| 187 |
+
)
|
| 188 |
|
| 189 |
# Apply token scaling if requested
|
| 190 |
if use_token_scaling:
|
|
|
|
| 249 |
def fused_linear_cross_entropy_backward(grad_output, grad_input, grad_weight, grad_bias):
|
| 250 |
# If cross entropy is the last layer, grad_output is 1.0. Skip the mul to save time
|
| 251 |
if not torch.equal(grad_output, torch.tensor(1.0, device=grad_output.device)):
|
| 252 |
+
with device_context(grad_input.device):
|
| 253 |
+
# We use a Triton kernel instead of a PyTorch operation because modifying inputs in-place
|
| 254 |
+
# for gradient storage and backward multiple times causes anomalies with PyTorch but not with Triton.
|
| 255 |
+
BT, H = grad_input.shape
|
| 256 |
+
n_rows = BT
|
| 257 |
+
BLOCK_SIZE = min(MAX_FUSED_SIZE, triton.next_power_of_2(H))
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 258 |
|
| 259 |
element_mul_kernel[(n_rows,)](
|
| 260 |
+
grad_input,
|
| 261 |
+
grad_input.stride(-2),
|
| 262 |
grad_output,
|
| 263 |
H,
|
| 264 |
BLOCK_SIZE=BLOCK_SIZE,
|
| 265 |
num_warps=32 if not is_hip() else 16,
|
| 266 |
)
|
| 267 |
|
| 268 |
+
# handle grad_weight
|
| 269 |
+
if grad_weight is not None:
|
| 270 |
+
V, H = grad_weight.shape
|
| 271 |
+
n_rows = V
|
| 272 |
+
|
| 273 |
+
element_mul_kernel[(n_rows,)](
|
| 274 |
+
grad_weight,
|
| 275 |
+
grad_weight.stride(-2),
|
| 276 |
+
grad_output,
|
| 277 |
+
H,
|
| 278 |
+
BLOCK_SIZE=BLOCK_SIZE,
|
| 279 |
+
num_warps=32 if not is_hip() else 16,
|
| 280 |
+
)
|
| 281 |
+
|
| 282 |
+
if grad_bias is not None:
|
| 283 |
+
V = grad_bias.shape[0]
|
| 284 |
+
n_rows = V
|
| 285 |
+
|
| 286 |
+
element_mul_kernel[(n_rows,)](
|
| 287 |
+
grad_bias,
|
| 288 |
+
grad_bias.stride(-1),
|
| 289 |
+
grad_output,
|
| 290 |
+
1,
|
| 291 |
+
BLOCK_SIZE=BLOCK_SIZE,
|
| 292 |
+
num_warps=32 if not is_hip() else 16,
|
| 293 |
+
)
|
| 294 |
return grad_input, grad_weight, grad_bias
|
| 295 |
|
| 296 |
|
build/torch-rocm/geglu.py
CHANGED
|
@@ -8,6 +8,8 @@ from .utils import calculate_settings
|
|
| 8 |
from .utils import compare_version
|
| 9 |
from .utils import ensure_contiguous
|
| 10 |
from .utils import is_npu_available
|
|
|
|
|
|
|
| 11 |
|
| 12 |
if compare_version("triton", operator.ge, "3.0.0") and not is_npu_available():
|
| 13 |
try:
|
|
@@ -94,15 +96,16 @@ def geglu_forward(a, b):
|
|
| 94 |
|
| 95 |
BLOCK_SIZE, num_warps = calculate_settings(n_cols)
|
| 96 |
|
| 97 |
-
|
| 98 |
-
|
| 99 |
-
|
| 100 |
-
|
| 101 |
-
|
| 102 |
-
|
| 103 |
-
|
| 104 |
-
|
| 105 |
-
|
|
|
|
| 106 |
return a, b, c.view(*ori_shape)
|
| 107 |
|
| 108 |
|
|
@@ -114,15 +117,16 @@ def geglu_backward(a, b, dc):
|
|
| 114 |
|
| 115 |
BLOCK_SIZE, num_warps = calculate_settings(n_cols)
|
| 116 |
|
| 117 |
-
|
| 118 |
-
|
| 119 |
-
|
| 120 |
-
|
| 121 |
-
|
| 122 |
-
|
| 123 |
-
|
| 124 |
-
|
| 125 |
-
|
|
|
|
| 126 |
|
| 127 |
return a.view(*ori_shape), b.view(*ori_shape)
|
| 128 |
|
|
|
|
| 8 |
from .utils import compare_version
|
| 9 |
from .utils import ensure_contiguous
|
| 10 |
from .utils import is_npu_available
|
| 11 |
+
from .utils import device_context
|
| 12 |
+
|
| 13 |
|
| 14 |
if compare_version("triton", operator.ge, "3.0.0") and not is_npu_available():
|
| 15 |
try:
|
|
|
|
| 96 |
|
| 97 |
BLOCK_SIZE, num_warps = calculate_settings(n_cols)
|
| 98 |
|
| 99 |
+
with device_context(a.device):
|
| 100 |
+
_geglu_tanh_forward_kernel[(n_rows,)](
|
| 101 |
+
a,
|
| 102 |
+
b,
|
| 103 |
+
c,
|
| 104 |
+
c.stride(-2),
|
| 105 |
+
n_cols=n_cols,
|
| 106 |
+
BLOCK_SIZE=BLOCK_SIZE,
|
| 107 |
+
num_warps=num_warps,
|
| 108 |
+
)
|
| 109 |
return a, b, c.view(*ori_shape)
|
| 110 |
|
| 111 |
|
|
|
|
| 117 |
|
| 118 |
BLOCK_SIZE, num_warps = calculate_settings(n_cols)
|
| 119 |
|
| 120 |
+
with device_context(a.device):
|
| 121 |
+
_geglu_tanh_backward_kernel[(n_rows,)](
|
| 122 |
+
dc,
|
| 123 |
+
a,
|
| 124 |
+
b,
|
| 125 |
+
dc.stride(-2),
|
| 126 |
+
n_cols=n_cols,
|
| 127 |
+
BLOCK_SIZE=BLOCK_SIZE,
|
| 128 |
+
num_warps=num_warps,
|
| 129 |
+
)
|
| 130 |
|
| 131 |
return a.view(*ori_shape), b.view(*ori_shape)
|
| 132 |
|
build/torch-rocm/group_norm.py
CHANGED
|
@@ -8,6 +8,8 @@ from .utils import compare_version
|
|
| 8 |
from .utils import ensure_contiguous
|
| 9 |
from .utils import infer_device
|
| 10 |
from .utils import is_npu_available
|
|
|
|
|
|
|
| 11 |
|
| 12 |
if compare_version("triton", operator.ge, "3.0.0") and not is_npu_available():
|
| 13 |
try:
|
|
@@ -215,26 +217,27 @@ def group_norm_forward(X, num_channels, num_groups, W, B, eps):
|
|
| 215 |
Mean = torch.zeros((batch_size, num_groups), dtype=X.dtype, device=X.device)
|
| 216 |
RSTD = torch.zeros((batch_size, num_groups), dtype=X.dtype, device=X.device)
|
| 217 |
|
| 218 |
-
|
| 219 |
-
|
| 220 |
-
|
| 221 |
-
|
| 222 |
-
|
| 223 |
-
|
| 224 |
-
|
| 225 |
-
|
| 226 |
-
|
| 227 |
-
|
| 228 |
-
|
| 229 |
-
|
| 230 |
-
|
| 231 |
-
|
| 232 |
-
|
| 233 |
-
|
| 234 |
-
|
| 235 |
-
|
| 236 |
-
|
| 237 |
-
|
|
|
|
| 238 |
# Return tensors in the original shape
|
| 239 |
return Y.view(*shape), X.view(*shape), Mean, RSTD, BLOCK_SIZE
|
| 240 |
|
|
@@ -254,25 +257,26 @@ def group_norm_backward(dY, X, W, B, Mean, RSTD, num_channels, num_groups):
|
|
| 254 |
DB = torch.zeros((num_channels), dtype=B.dtype, device=B.device)
|
| 255 |
triton_dtype = tl.float32 if X.dtype == torch.float32 else tl.bfloat16
|
| 256 |
|
| 257 |
-
|
| 258 |
-
|
| 259 |
-
|
| 260 |
-
|
| 261 |
-
|
| 262 |
-
|
| 263 |
-
|
| 264 |
-
|
| 265 |
-
|
| 266 |
-
|
| 267 |
-
|
| 268 |
-
|
| 269 |
-
|
| 270 |
-
|
| 271 |
-
|
| 272 |
-
|
| 273 |
-
|
| 274 |
-
|
| 275 |
-
|
|
|
|
| 276 |
|
| 277 |
# Return tensors in the original shape
|
| 278 |
return DX.view(*shape), DW, DB
|
|
|
|
| 8 |
from .utils import ensure_contiguous
|
| 9 |
from .utils import infer_device
|
| 10 |
from .utils import is_npu_available
|
| 11 |
+
from .utils import device_context
|
| 12 |
+
|
| 13 |
|
| 14 |
if compare_version("triton", operator.ge, "3.0.0") and not is_npu_available():
|
| 15 |
try:
|
|
|
|
| 217 |
Mean = torch.zeros((batch_size, num_groups), dtype=X.dtype, device=X.device)
|
| 218 |
RSTD = torch.zeros((batch_size, num_groups), dtype=X.dtype, device=X.device)
|
| 219 |
|
| 220 |
+
with device_context(X.device):
|
| 221 |
+
_group_norm_forward_kernel[(batch_size, num_groups)](
|
| 222 |
+
Y,
|
| 223 |
+
Y.stride(0),
|
| 224 |
+
Y.stride(1),
|
| 225 |
+
X,
|
| 226 |
+
X.stride(0),
|
| 227 |
+
X.stride(1),
|
| 228 |
+
Mean,
|
| 229 |
+
Mean.stride(0),
|
| 230 |
+
Mean.stride(1),
|
| 231 |
+
RSTD,
|
| 232 |
+
RSTD.stride(0),
|
| 233 |
+
RSTD.stride(1),
|
| 234 |
+
W,
|
| 235 |
+
B,
|
| 236 |
+
hidden_size,
|
| 237 |
+
channels_per_group,
|
| 238 |
+
eps,
|
| 239 |
+
BLOCK_SIZE=BLOCK_SIZE,
|
| 240 |
+
)
|
| 241 |
# Return tensors in the original shape
|
| 242 |
return Y.view(*shape), X.view(*shape), Mean, RSTD, BLOCK_SIZE
|
| 243 |
|
|
|
|
| 257 |
DB = torch.zeros((num_channels), dtype=B.dtype, device=B.device)
|
| 258 |
triton_dtype = tl.float32 if X.dtype == torch.float32 else tl.bfloat16
|
| 259 |
|
| 260 |
+
with device_context(X.device):
|
| 261 |
+
BLOCK_SIZE = min(MAX_FUSED_SIZE, triton.next_power_of_2(hidden_size))
|
| 262 |
+
_group_norm_backward_kernel[(batch_size, num_groups)](
|
| 263 |
+
X,
|
| 264 |
+
X.stride(0),
|
| 265 |
+
X.stride(1),
|
| 266 |
+
W,
|
| 267 |
+
Mean,
|
| 268 |
+
Mean.stride(0),
|
| 269 |
+
Mean.stride(1),
|
| 270 |
+
RSTD,
|
| 271 |
+
DX,
|
| 272 |
+
DW,
|
| 273 |
+
DB,
|
| 274 |
+
dY,
|
| 275 |
+
hidden_size,
|
| 276 |
+
channels_per_group,
|
| 277 |
+
BLOCK_SIZE=BLOCK_SIZE,
|
| 278 |
+
dtype=triton_dtype,
|
| 279 |
+
)
|
| 280 |
|
| 281 |
# Return tensors in the original shape
|
| 282 |
return DX.view(*shape), DW, DB
|
build/torch-rocm/jsd.py
CHANGED
|
@@ -6,6 +6,7 @@ import triton.language as tl
|
|
| 6 |
|
| 7 |
from .utils import ensure_contiguous
|
| 8 |
from .utils import infer_device
|
|
|
|
| 9 |
|
| 10 |
|
| 11 |
@triton.jit
|
|
@@ -109,23 +110,24 @@ def jsd_forward(_input, target, shift_labels, beta, ignore_index, has_label):
|
|
| 109 |
else:
|
| 110 |
n_non_ignore = BT
|
| 111 |
|
| 112 |
-
|
| 113 |
-
|
| 114 |
-
|
| 115 |
-
|
| 116 |
-
|
| 117 |
-
|
| 118 |
-
|
| 119 |
-
|
| 120 |
-
|
| 121 |
-
|
| 122 |
-
|
| 123 |
-
|
| 124 |
-
|
| 125 |
-
|
| 126 |
-
|
| 127 |
-
|
| 128 |
-
|
|
|
|
| 129 |
|
| 130 |
loss = torch.sum(loss)
|
| 131 |
return loss.to(_input.dtype), dX
|
|
|
|
| 6 |
|
| 7 |
from .utils import ensure_contiguous
|
| 8 |
from .utils import infer_device
|
| 9 |
+
from .utils import device_context
|
| 10 |
|
| 11 |
|
| 12 |
@triton.jit
|
|
|
|
| 110 |
else:
|
| 111 |
n_non_ignore = BT
|
| 112 |
|
| 113 |
+
with device_context(_input.device):
|
| 114 |
+
_jsd_kernel[(n_rows,)](
|
| 115 |
+
X_ptr=_input, # input in logspace, X = log Q
|
| 116 |
+
X_stride=_input.stride(-2),
|
| 117 |
+
Y_ptr=target, # ground truth in logspace, Y = log P
|
| 118 |
+
Y_stride=target.stride(-2),
|
| 119 |
+
loss_ptr=loss,
|
| 120 |
+
loss_stride=loss.stride(-2),
|
| 121 |
+
dX_ptr=dX,
|
| 122 |
+
dX_stride=dX.stride(-2),
|
| 123 |
+
label_ptr=(shift_labels if has_label else torch.empty(1, device=_input.device)), # dummy ptr if no label
|
| 124 |
+
beta=beta,
|
| 125 |
+
n_non_ignore=n_non_ignore,
|
| 126 |
+
ignore_index=ignore_index,
|
| 127 |
+
n_cols=V,
|
| 128 |
+
BLOCK_SIZE=BLOCK_SIZE,
|
| 129 |
+
HAS_LABEL=has_label,
|
| 130 |
+
)
|
| 131 |
|
| 132 |
loss = torch.sum(loss)
|
| 133 |
return loss.to(_input.dtype), dX
|
build/torch-rocm/kl_div.py
CHANGED
|
@@ -7,6 +7,7 @@ import triton.language as tl
|
|
| 7 |
from .utils import ensure_contiguous
|
| 8 |
from .utils import is_hip
|
| 9 |
from .utils import infer_device
|
|
|
|
| 10 |
|
| 11 |
|
| 12 |
def get_num_warps(BLOCK_SIZE):
|
|
@@ -130,20 +131,21 @@ def kldiv_forward_triton(y_pred, y_true, log_target, reduction, eps): # [BT, V]
|
|
| 130 |
out_size = (BT, V) if reduction == _REDUCTION_MODE_NONE.value else (BT,)
|
| 131 |
output_tensor = torch.zeros(out_size, device=y_pred.device, dtype=torch.float32)
|
| 132 |
|
| 133 |
-
|
| 134 |
-
|
| 135 |
-
|
| 136 |
-
|
| 137 |
-
|
| 138 |
-
|
| 139 |
-
|
| 140 |
-
|
| 141 |
-
|
| 142 |
-
|
| 143 |
-
|
| 144 |
-
|
| 145 |
-
|
| 146 |
-
|
|
|
|
| 147 |
|
| 148 |
# calculated according to the reduction mode same as in Pytorch. In the later versions, `mean` will be changed to the same behavior as `batchmean`
|
| 149 |
# https://pytorch.org/docs/stable/generated/torch.nn.KLDivLoss.html
|
|
@@ -165,17 +167,18 @@ def kldiv_backward_triton(target, grad_output, new_grads, log_target):
|
|
| 165 |
|
| 166 |
grid = (BT,)
|
| 167 |
|
| 168 |
-
|
| 169 |
-
|
| 170 |
-
|
| 171 |
-
|
| 172 |
-
|
| 173 |
-
|
| 174 |
-
|
| 175 |
-
|
| 176 |
-
|
| 177 |
-
|
| 178 |
-
|
|
|
|
| 179 |
|
| 180 |
# If cross entropy is the last layer, grad_output is 1.0. Skip the mul then.
|
| 181 |
if torch.equal(grad_output, torch.tensor(1.0, device=grad_output.device)):
|
|
|
|
| 7 |
from .utils import ensure_contiguous
|
| 8 |
from .utils import is_hip
|
| 9 |
from .utils import infer_device
|
| 10 |
+
from .utils import device_context
|
| 11 |
|
| 12 |
|
| 13 |
def get_num_warps(BLOCK_SIZE):
|
|
|
|
| 131 |
out_size = (BT, V) if reduction == _REDUCTION_MODE_NONE.value else (BT,)
|
| 132 |
output_tensor = torch.zeros(out_size, device=y_pred.device, dtype=torch.float32)
|
| 133 |
|
| 134 |
+
with device_context(y_pred.device):
|
| 135 |
+
_kldiv_kernel_forward[grid](
|
| 136 |
+
y_pred,
|
| 137 |
+
y_pred.stride(0),
|
| 138 |
+
y_true,
|
| 139 |
+
y_true.stride(0),
|
| 140 |
+
output_tensor,
|
| 141 |
+
output_tensor.stride(0),
|
| 142 |
+
V,
|
| 143 |
+
eps=eps,
|
| 144 |
+
BLOCK_SIZE=BLOCK_SIZE,
|
| 145 |
+
num_warps=num_warps,
|
| 146 |
+
log_target=log_target,
|
| 147 |
+
reduction=reduction,
|
| 148 |
+
)
|
| 149 |
|
| 150 |
# calculated according to the reduction mode same as in Pytorch. In the later versions, `mean` will be changed to the same behavior as `batchmean`
|
| 151 |
# https://pytorch.org/docs/stable/generated/torch.nn.KLDivLoss.html
|
|
|
|
| 167 |
|
| 168 |
grid = (BT,)
|
| 169 |
|
| 170 |
+
with device_context(target.device):
|
| 171 |
+
# We store the gradients in-place in the input tensor
|
| 172 |
+
_kldiv_kernel_backward[grid](
|
| 173 |
+
target,
|
| 174 |
+
target.stride(0),
|
| 175 |
+
new_grads,
|
| 176 |
+
new_grads.stride(0),
|
| 177 |
+
V,
|
| 178 |
+
BLOCK_SIZE=BLOCK_SIZE,
|
| 179 |
+
num_warps=num_warps,
|
| 180 |
+
log_target=log_target,
|
| 181 |
+
)
|
| 182 |
|
| 183 |
# If cross entropy is the last layer, grad_output is 1.0. Skip the mul then.
|
| 184 |
if torch.equal(grad_output, torch.tensor(1.0, device=grad_output.device)):
|
build/torch-rocm/layer_norm.py
CHANGED
|
@@ -11,6 +11,8 @@ from .utils import ensure_contiguous
|
|
| 11 |
from .utils import get_npu_core_count
|
| 12 |
from .utils import set_large_grf_mode
|
| 13 |
from .utils import is_npu_available
|
|
|
|
|
|
|
| 14 |
|
| 15 |
if compare_version("triton", operator.ge, "3.0.0") and not is_npu_available():
|
| 16 |
try:
|
|
@@ -202,26 +204,27 @@ def layer_norm_forward(X, W, B, eps):
|
|
| 202 |
if X.device.type == "xpu":
|
| 203 |
set_large_grf_mode(kernel_args)
|
| 204 |
|
| 205 |
-
|
| 206 |
-
|
| 207 |
-
|
| 208 |
-
|
| 209 |
-
|
| 210 |
-
|
| 211 |
-
|
| 212 |
-
|
| 213 |
-
|
| 214 |
-
|
| 215 |
-
|
| 216 |
-
|
| 217 |
-
|
| 218 |
-
|
| 219 |
-
|
| 220 |
-
|
| 221 |
-
|
| 222 |
-
|
| 223 |
-
|
| 224 |
-
|
|
|
|
| 225 |
)
|
| 226 |
|
| 227 |
return Y.view(*shape), X, Mean, RSTD, BLOCK_SIZE, num_warps
|
|
@@ -273,29 +276,30 @@ def layer_norm_backward(dY, X, W, B, Mean, RSTD):
|
|
| 273 |
kernel_args.update({"num_warps": 32, "num_stages": 4})
|
| 274 |
set_large_grf_mode(kernel_args)
|
| 275 |
|
| 276 |
-
|
| 277 |
-
|
| 278 |
-
|
| 279 |
-
|
| 280 |
-
|
| 281 |
-
|
| 282 |
-
|
| 283 |
-
|
| 284 |
-
|
| 285 |
-
|
| 286 |
-
|
| 287 |
-
|
| 288 |
-
|
| 289 |
-
|
| 290 |
-
|
| 291 |
-
|
| 292 |
-
|
| 293 |
-
|
| 294 |
-
|
| 295 |
-
|
| 296 |
-
|
| 297 |
-
|
| 298 |
-
|
|
|
|
| 299 |
|
| 300 |
DX = DX.view(*shape)
|
| 301 |
DW = _DW.sum(dim=0).to(W.dtype)
|
|
|
|
| 11 |
from .utils import get_npu_core_count
|
| 12 |
from .utils import set_large_grf_mode
|
| 13 |
from .utils import is_npu_available
|
| 14 |
+
from .utils import device_context
|
| 15 |
+
|
| 16 |
|
| 17 |
if compare_version("triton", operator.ge, "3.0.0") and not is_npu_available():
|
| 18 |
try:
|
|
|
|
| 204 |
if X.device.type == "xpu":
|
| 205 |
set_large_grf_mode(kernel_args)
|
| 206 |
|
| 207 |
+
with device_context(X.device):
|
| 208 |
+
# Launch kernel with one thread block per row for optimal performance
|
| 209 |
+
grid = (n_rows,)
|
| 210 |
+
_layer_norm_forward_kernel[grid](
|
| 211 |
+
Y,
|
| 212 |
+
Y.stride(0),
|
| 213 |
+
X,
|
| 214 |
+
X.stride(0),
|
| 215 |
+
W,
|
| 216 |
+
W.stride(0),
|
| 217 |
+
B,
|
| 218 |
+
B.stride(0),
|
| 219 |
+
Mean,
|
| 220 |
+
Mean.stride(0),
|
| 221 |
+
RSTD,
|
| 222 |
+
RSTD.stride(0),
|
| 223 |
+
n_cols,
|
| 224 |
+
eps,
|
| 225 |
+
BLOCK_SIZE=BLOCK_SIZE,
|
| 226 |
+
num_warps=num_warps,
|
| 227 |
+
**kernel_args,
|
| 228 |
)
|
| 229 |
|
| 230 |
return Y.view(*shape), X, Mean, RSTD, BLOCK_SIZE, num_warps
|
|
|
|
| 276 |
kernel_args.update({"num_warps": 32, "num_stages": 4})
|
| 277 |
set_large_grf_mode(kernel_args)
|
| 278 |
|
| 279 |
+
with device_context(X.device):
|
| 280 |
+
# Launch kernel with one thread block per row for optimal performance
|
| 281 |
+
_layer_norm_backward_kernel[grid](
|
| 282 |
+
X,
|
| 283 |
+
X.stride(0),
|
| 284 |
+
W,
|
| 285 |
+
Mean,
|
| 286 |
+
Mean.stride(0),
|
| 287 |
+
RSTD,
|
| 288 |
+
RSTD.stride(0),
|
| 289 |
+
DX,
|
| 290 |
+
DX.stride(0),
|
| 291 |
+
_DW,
|
| 292 |
+
_DW.stride(0),
|
| 293 |
+
_DB,
|
| 294 |
+
_DB.stride(0),
|
| 295 |
+
dY,
|
| 296 |
+
dY.stride(0),
|
| 297 |
+
n_rows,
|
| 298 |
+
n_cols,
|
| 299 |
+
rows_per_program=rows_per_program,
|
| 300 |
+
BLOCK_SIZE=BLOCK_SIZE,
|
| 301 |
+
**kernel_args,
|
| 302 |
+
)
|
| 303 |
|
| 304 |
DX = DX.view(*shape)
|
| 305 |
DW = _DW.sum(dim=0).to(W.dtype)
|
build/torch-rocm/metadata.json
CHANGED
|
@@ -1,6 +1,6 @@
|
|
| 1 |
{
|
| 2 |
"name": "liger-kernels",
|
| 3 |
-
"id": "
|
| 4 |
"version": 3,
|
| 5 |
"license": "BSD-2-Clause",
|
| 6 |
"python-depends": [],
|
|
@@ -11,24 +11,24 @@
|
|
| 11 |
"algorithm": "sha256",
|
| 12 |
"files": {
|
| 13 |
"__init__.py": "DSZMiK0xOBiMb0JdCc2K9QH4Z6lV7L8D/fTowstf/og=",
|
| 14 |
-
"_ops.py": "
|
| 15 |
-
"cross_entropy.py": "
|
| 16 |
-
"dyt.py": "
|
| 17 |
-
"fused_linear_cross_entropy.py": "
|
| 18 |
-
"geglu.py": "
|
| 19 |
-
"group_norm.py": "
|
| 20 |
-
"jsd.py": "
|
| 21 |
-
"kl_div.py": "
|
| 22 |
-
"layer_norm.py": "
|
| 23 |
"layers.py": "+50F9xmnpQXbU1K/lZxdJwbnTotkksJlhoLOWHjo+/M=",
|
| 24 |
"liger_kernels/__init__.py": "DFYPlrhXwYjEqCl/8n0SmWGZV8NFml5DPhMjKfv98GY=",
|
| 25 |
-
"qwen2vl_mrope.py": "
|
| 26 |
-
"rms_norm.py": "
|
| 27 |
-
"rope.py": "
|
| 28 |
-
"swiglu.py": "
|
| 29 |
"tiled_mlp.py": "PJw2R8YdHHxQRSwbFJWjcHxgzssEqsRWqyogkY8pKR8=",
|
| 30 |
-
"tvd.py": "
|
| 31 |
-
"utils.py": "
|
| 32 |
}
|
| 33 |
},
|
| 34 |
"provenance": {
|
|
@@ -38,7 +38,7 @@
|
|
| 38 |
"dirty": false
|
| 39 |
},
|
| 40 |
"kernel": {
|
| 41 |
-
"sha": "
|
| 42 |
"dirty": false
|
| 43 |
}
|
| 44 |
}
|
|
|
|
| 1 |
{
|
| 2 |
"name": "liger-kernels",
|
| 3 |
+
"id": "_liger_kernels_rocm_0c5fb33",
|
| 4 |
"version": 3,
|
| 5 |
"license": "BSD-2-Clause",
|
| 6 |
"python-depends": [],
|
|
|
|
| 11 |
"algorithm": "sha256",
|
| 12 |
"files": {
|
| 13 |
"__init__.py": "DSZMiK0xOBiMb0JdCc2K9QH4Z6lV7L8D/fTowstf/og=",
|
| 14 |
+
"_ops.py": "Z4DkRS77DKP3IQGbv+oQOt0NNFzokK2hpXfcHoNpG/I=",
|
| 15 |
+
"cross_entropy.py": "N/GPXvZTWjsu2qlAC96kuFJ4XFI0c3Qff9YtOkqaQG0=",
|
| 16 |
+
"dyt.py": "Fa7TEHZqkfNKXaxSYQt8/MlnOUFo+WIUIkz6ZA2O/b8=",
|
| 17 |
+
"fused_linear_cross_entropy.py": "l21OYH+10ryswTwtuad+h7V6Xo/8zeSLoKDyn7ZCVxo=",
|
| 18 |
+
"geglu.py": "m1Y4vqkXPBdoSzxKOTi7SiU8NW85a2JW/2/2BQbgktg=",
|
| 19 |
+
"group_norm.py": "7EcdehRRoGaias7RiU5Qil+AYQ8ZClP7mJosl1Og8r4=",
|
| 20 |
+
"jsd.py": "yWhyIa06nzvlhcwGBQNSGrTv4bessgm4vrgqQEfNUMo=",
|
| 21 |
+
"kl_div.py": "nPN5Lb2NcIGWzk+RNxgFJ69GKyfPeRLg4ycCns7KO9U=",
|
| 22 |
+
"layer_norm.py": "+t987/DupvUHRaLK8s5V9W6ewHtFZIG5WJLD5kg0RjA=",
|
| 23 |
"layers.py": "+50F9xmnpQXbU1K/lZxdJwbnTotkksJlhoLOWHjo+/M=",
|
| 24 |
"liger_kernels/__init__.py": "DFYPlrhXwYjEqCl/8n0SmWGZV8NFml5DPhMjKfv98GY=",
|
| 25 |
+
"qwen2vl_mrope.py": "uxztQJ0LlMEP4Zec/yjxVJa0/uOsB5XvjUN5rM37Rf4=",
|
| 26 |
+
"rms_norm.py": "LJfKK/kutmbztb17eC3a6OGAk39PLEAHa+hzzmcCjAA=",
|
| 27 |
+
"rope.py": "IdrugShnaNJJ9apTXlMEXMjZE0IpDZI7X23gIgYS3k4=",
|
| 28 |
+
"swiglu.py": "QSjLVgdF3a4rsPdI0B0VNTmA26M2F4Y9SSv4+303WkE=",
|
| 29 |
"tiled_mlp.py": "PJw2R8YdHHxQRSwbFJWjcHxgzssEqsRWqyogkY8pKR8=",
|
| 30 |
+
"tvd.py": "BEL3OibZ2tfGJRyycMTSUWah/6eDdLSV2jA7ZKcxE7s=",
|
| 31 |
+
"utils.py": "l4/OveROs0DFMcRaJsaoiJBKq7jvHbAL9aRRrkZb+iM="
|
| 32 |
}
|
| 33 |
},
|
| 34 |
"provenance": {
|
|
|
|
| 38 |
"dirty": false
|
| 39 |
},
|
| 40 |
"kernel": {
|
| 41 |
+
"sha": "0c5fb33d21084e07750a248da12071ff011bd70a",
|
| 42 |
"dirty": false
|
| 43 |
}
|
| 44 |
}
|
build/torch-rocm/metadata.json.sigstore
CHANGED
|
@@ -1 +1 @@
|
|
| 1 |
-
{"mediaType":"application/vnd.dev.sigstore.bundle.v0.3+json", "verificationMaterial":{"certificate":{"rawBytes":"MIIHTDCCBtGgAwIBAgIUeKKjysjZHPTkcjmFb3CHS1TysN0wCgYIKoZIzj0EAwMwNzEVMBMGA1UEChMMc2lnc3RvcmUuZGV2MR4wHAYDVQQDExVzaWdzdG9yZS1pbnRlcm1lZGlhdGUwHhcNMjYwNzE0MTIwMjUwWhcNMjYwNzE0MTIxMjUwWjAAMFkwEwYHKoZIzj0CAQYIKoZIzj0DAQcDQgAEMw0GrkZBBP9oqnpDe5wf8q41n6WefaLaklEXON/6Xd4W4HwqUuA5A2sreVnQq1EjKZd4nZvn0yIbO+G/5kGEfKOCBfAwggXsMA4GA1UdDwEB/wQEAwIHgDATBgNVHSUEDDAKBggrBgEFBQcDAzAdBgNVHQ4EFgQUtZx+9iAL0iQdwe1N8MZBeaC2LFIwHwYDVR0jBBgwFoAU39Ppz1YkEZb5qNjpKFWixi4YZD8wawYDVR0RAQH/BGEwX4ZdaHR0cHM6Ly9naXRodWIuY29tL2h1Z2dpbmdmYWNlL2tlcm5lbHMtY29tbXVuaXR5Ly5naXRodWIvd29ya2Zsb3dzL2J1aWxkLnlhbWxAcmVmcy9oZWFkcy9tYWluMDkGCisGAQQBg78wAQEEK2h0dHBzOi8vdG9rZW4uYWN0aW9ucy5naXRodWJ1c2VyY29udGVudC5jb20wHwYKKwYBBAGDvzABAgQRd29ya2Zsb3dfZGlzcGF0Y2gwNgYKKwYBBAGDvzABAwQoNGQ5Zjc5OGI5NWVkMjgwYjc3MDIxMzA0OGI4ZDNmNDE1MDlkMGM5ZTATBgorBgEEAYO/MAEEBAVCdWlsZDArBgorBgEEAYO/MAEFBB1odWdnaW5nZmFjZS9rZXJuZWxzLWNvbW11bml0eTAdBgorBgEEAYO/MAEGBA9yZWZzL2hlYWRzL21haW4wOwYKKwYBBAGDvzABCAQtDCtodHRwczovL3Rva2VuLmFjdGlvbnMuZ2l0aHVidXNlcmNvbnRlbnQuY29tMG0GCisGAQQBg78wAQkEXwxdaHR0cHM6Ly9naXRodWIuY29tL2h1Z2dpbmdmYWNlL2tlcm5lbHMtY29tbXVuaXR5Ly5naXRodWIvd29ya2Zsb3dzL2J1aWxkLnlhbWxAcmVmcy9oZWFkcy9tYWluMDgGCisGAQQBg78wAQoEKgwoNGQ5Zjc5OGI5NWVkMjgwYjc3MDIxMzA0OGI4ZDNmNDE1MDlkMGM5ZTAbBgorBgEEAYO/MAELBA0MC3NlbGYtaG9zdGVkMEAGCisGAQQBg78wAQwEMgwwaHR0cHM6Ly9naXRodWIuY29tL2h1Z2dpbmdmYWNlL2tlcm5lbHMtY29tbXVuaXR5MDgGCisGAQQBg78wAQ0EKgwoNGQ5Zjc5OGI5NWVkMjgwYjc3MDIxMzA0OGI4ZDNmNDE1MDlkMGM5ZTAfBgorBgEEAYO/MAEOBBEMD3JlZnMvaGVhZHMvbWFpbjAaBgorBgEEAYO/MAEPBAwMCjEwNzE0NzU1MjkwLgYKKwYBBAGDvzABEAQgDB5odHRwczovL2dpdGh1Yi5jb20vaHVnZ2luZ2ZhY2UwGAYKKwYBBAGDvzABEQQKDAgyNTcyMDc0MzBtBgorBgEEAYO/MAESBF8MXWh0dHBzOi8vZ2l0aHViLmNvbS9odWdnaW5nZmFjZS9rZXJuZWxzLWNvbW11bml0eS8uZ2l0aHViL3dvcmtmbG93cy9idWlsZC55YW1sQHJlZnMvaGVhZHMvbWFpbjA4BgorBgEEAYO/MAETBCoMKDRkOWY3OThiOTVlZDI4MGI3NzAyMTMwNDhiOGQzZjQxNTA5ZDBjOWUwIQYKKwYBBAGDvzABFAQTDBF3b3JrZmxvd19kaXNwYXRjaDBkBgorBgEEAYO/MAEVBFYMVGh0dHBzOi8vZ2l0aHViLmNvbS9odWdnaW5nZmFjZS9rZXJuZWxzLWNvbW11bml0eS9hY3Rpb25zL3J1bnMvMjkzMzA2MjM4NjUvYXR0ZW1wdHMvMTAWBgorBgEEAYO/MAEWBAgMBnB1YmxpYzBGBgorBgEEAYO/MAEYBDgMNnJlcG86aHVnZ2luZ2ZhY2Uva2VybmVscy1jb21tdW5pdHk6cmVmOnJlZnMvaGVhZHMvbWFpbjCBigYKKwYBBAHWeQIEAgR8BHoAeAB2AN09MGrGxxEyYxkeHJlnNwKiSl643jyt/4eKcoAvKe6OAAABn2CCKbEAAAQDAEcwRQIgMoVU3dtBnGlxMMb3sOTM8i/rOdiqaQ37x0tz9NOZa+MCIQCvrHRw6fB5k7rokMe6iDv+n9siBeeEDI5Helc4t8PNfzAKBggqhkjOPQQDAwNpADBmAjEAuazcBiTcv8K6MyCbxJkO8Q1WX752+NAYF6uQhREhOlfLY1Htf2AgU5YzxWwv72nSAjEAsyIagTrwrpp0hqNOQI9ik7Cl4rMqprpQ90wKXEz5+urqVmetTmkd7YsMArDzriDo"}, "tlogEntries":[{"logIndex":"2167855782", "logId":{"keyId":"wNI9atQGlz+VWfO6LRygH4QUfY/8W4RFwiT5i5WRgB0="}, "kindVersion":{"kind":"hashedrekord", "version":"0.0.1"}, "integratedTime":"1784030571", "inclusionPromise":{"signedEntryTimestamp":"MEQCIFPwhR+59GwljQ+ZxD1rhaNTV3Sh4vPwQeU/9JrQoO2xAiBPUFa1Dck8XvuXhEpd/t74KU9P4qaGczgDzl//4hAk8w=="}, "inclusionProof":{"logIndex":"2045951520", "rootHash":"Re/Ph87BOM7jhvV4VE+3WsyQb72tumWdlqTxAOe0qEM=", "treeSize":"2045951527", "hashes":["nUla+MAuN3xIPh3Fw90o+gJvejDEBaNrSVpncG6U3kM=", "dEagCZIabgWBH6ViR5uwn8xXp3WdvaR3vEy+Q8US9b8=", "4tMZ/2/4pBxGx11gMAczICVOHyul021sgz6LWWWRpIQ=", "m+Vk8U0/4EcekuSwkXtYOJcYVvuBCDXI7XChKnMHSec=", "iMBom1ahASM6ZAxymZoy+ISJTz5W3t4BNlJ+0yiYrnc=", "IuHNjKNGzZgw8Ec1vaQ4yxMC8m5mfRHL/sO7dTeEGIM=", "T/vXn6bo+SCijMlRZwcNRxHpWcvInutRlkDvk1Y7oEM=", "TXWJIUMj+J24/w5Nh86FuokZClxMAnvB6sia+ArBfNI=", "pfPn/gRvLHfxa/LNybYhqG78qHZ4QD3PI+zQ35oYweM=", "sMaCrs8erYNQ48jT2FC56F0RAbpTkuSXZ+cs1zcSObw=", "zGam2EiRYwS9IjemppqM/acVfU15hw3eAqn9Fb8+8vY=", "qQGtNWYnFxNQ1OmqNBR2y92EFrgo5lyaKNQLaE/Dl1c=", "/UAkaaq9D8BPMw5qLAxaf+ur4Sfbv3FfLWB2M/py3wM=", "56+1wf4OgrBZEzfW9n13hPfm0bRvMRgerPv4TPE+3Ho=", "mTSsSZdzVhRgljy9CvDq3GfjxxeOTHrwgpOWK4rRj3Q=", "5sxNZDoxEj6DMmQATisX9bQXdFmeRHYfC8BjyDIIgec=", "Rq/A2aTC9e54ldNjcpsJ26rX/h8JlN6ZHxgw7yEa6C0=", "+/VZ56MsIPxMiyLAodzKXo5TEWdQp36z89qLhpzloAo=", "daxmZaajRpZV+JxHiOYZhJBiSKN5ucqjh2WnGbHhirw=", "DOCeoSMovIvLExkhIvisow9AuNXgeWs4ECkyR6EcqYU="], "checkpoint":{"envelope":"rekor.sigstore.dev - 1193050959916656506\n2045951527\nRe/Ph87BOM7jhvV4VE+3WsyQb72tumWdlqTxAOe0qEM=\n\n— rekor.sigstore.dev wNI9ajBEAiAxviZPcApj0ZJ1BUlQmvWCHwh/I05iSQg0PsjajUUu5wIgfTmpusfaD0Tu2YE2XuUcR3X1uFqLW7goWine86qQDnQ=\n"}}, "canonicalizedBody":"eyJhcGlWZXJzaW9uIjoiMC4wLjEiLCJraW5kIjoiaGFzaGVkcmVrb3JkIiwic3BlYyI6eyJkYXRhIjp7Imhhc2giOnsiYWxnb3JpdGhtIjoic2hhMjU2IiwidmFsdWUiOiJiNmJkMTIwYTVhOGNkYWE0MmI5MTkxYWYwYzBlZDQzNDMxMDE5OWU2Y2RkZDg1YzFhNTE4NTAxODE3Yzc2NDliIn19LCJzaWduYXR1cmUiOnsiY29udGVudCI6Ik1FWUNJUUNHMmE1Y2pEL09zY3UrdlpxWjVnbXB6S2xVTGhXZHhyRmVSSE00dlRJa0p3SWhBSyttWkE1WlBZUGwrZnJUSFdLYk5TZUxFaFh5TjdzZ3VCNmpMMWxBTEhHTCIsInB1YmxpY0tleSI6eyJjb250ZW50IjoiTFMwdExTMUNSVWRKVGlCRFJWSlVTVVpKUTBGVVJTMHRMUzB0Q2sxSlNVaFVSRU5EUW5SSFowRjNTVUpCWjBsVlpVdExhbmx6YWxwSVVGUnJZMnB0Um1JelEwaFRNVlI1YzA0d2QwTm5XVWxMYjFwSmVtb3dSVUYzVFhjS1RucEZWazFDVFVkQk1WVkZRMmhOVFdNeWJHNWpNMUoyWTIxVmRWcEhWakpOVWpSM1NFRlpSRlpSVVVSRmVGWjZZVmRrZW1SSE9YbGFVekZ3WW01U2JBcGpiVEZzV2tkc2FHUkhWWGRJYUdOT1RXcFpkMDU2UlRCTlZFbDNUV3BWZDFkb1kwNU5hbGwzVG5wRk1FMVVTWGhOYWxWM1YycEJRVTFHYTNkRmQxbElDa3R2V2tsNmFqQkRRVkZaU1V0dldrbDZhakJFUVZGalJGRm5RVVZOZHpCSGNtdGFRa0pRT1c5eGJuQkVaVFYzWmpoeE5ERnVObGRsWm1GTVlXdHNSVmdLVDA0dk5saGtORmMwU0hkeFZYVkJOVUV5YzNKbFZtNVJjVEZGYWt0YVpEUnVXblp1TUhsSllrOHJSeTgxYTBkRlprdFBRMEptUVhkbloxaHpUVUUwUndwQk1WVmtSSGRGUWk5M1VVVkJkMGxJWjBSQlZFSm5UbFpJVTFWRlJFUkJTMEpuWjNKQ1owVkdRbEZqUkVGNlFXUkNaMDVXU0ZFMFJVWm5VVlYwV25nckNqbHBRVXd3YVZGa2QyVXhUamhOV2tKbFlVTXlURVpKZDBoM1dVUldVakJxUWtKbmQwWnZRVlV6T1ZCd2VqRlphMFZhWWpWeFRtcHdTMFpYYVhocE5Ga0tXa1E0ZDJGM1dVUldVakJTUVZGSUwwSkhSWGRZTkZwa1lVaFNNR05JVFRaTWVUbHVZVmhTYjJSWFNYVlpNamwwVERKb01Wb3laSEJpYldSdFdWZE9iQXBNTW5Sc1kyMDFiR0pJVFhSWk1qbDBZbGhXZFdGWVVqVk1lVFZ1WVZoU2IyUlhTWFprTWpsNVlUSmFjMkl6WkhwTU1rb3hZVmQ0YTB4dWJHaGlWM2hCQ21OdFZtMWplVGx2V2xkR2EyTjVPWFJaVjJ4MVRVUnJSME5wYzBkQlVWRkNaemM0ZDBGUlJVVkxNbWd3WkVoQ2VrOXBPSFprUnpseVdsYzBkVmxYVGpBS1lWYzVkV041Tlc1aFdGSnZaRmRLTVdNeVZubFpNamwxWkVkV2RXUkROV3BpTWpCM1NIZFpTMHQzV1VKQ1FVZEVkbnBCUWtGblVWSmtNamw1WVRKYWN3cGlNMlJtV2tkc2VtTkhSakJaTW1kM1RtZFpTMHQzV1VKQ1FVZEVkbnBCUWtGM1VXOU9SMUUxV21wak5VOUhTVFZPVjFaclRXcG5kMWxxWXpOTlJFbDRDazE2UVRCUFIwazBXa1JPYlU1RVJURk5SR3hyVFVkTk5WcFVRVlJDWjI5eVFtZEZSVUZaVHk5TlFVVkZRa0ZXUTJSWGJITmFSRUZ5UW1kdmNrSm5SVVVLUVZsUEwwMUJSVVpDUWpGdlpGZGtibUZYTlc1YWJVWnFXbE01Y2xwWVNuVmFWM2g2VEZkT2RtSlhNVEZpYld3d1pWUkJaRUpuYjNKQ1owVkZRVmxQTHdwTlFVVkhRa0U1ZVZwWFducE1NbWhzV1ZkU2Vrd3lNV2hoVnpSM1QzZFpTMHQzV1VKQ1FVZEVkbnBCUWtOQlVYUkVRM1J2WkVoU2QyTjZiM1pNTTFKMkNtRXlWblZNYlVacVpFZHNkbUp1VFhWYU1td3dZVWhXYVdSWVRteGpiVTUyWW01U2JHSnVVWFZaTWpsMFRVY3dSME5wYzBkQlVWRkNaemM0ZDBGUmEwVUtXSGQ0WkdGSVVqQmpTRTAyVEhrNWJtRllVbTlrVjBsMVdUSTVkRXd5YURGYU1tUndZbTFrYlZsWFRteE1NblJzWTIwMWJHSklUWFJaTWpsMFlsaFdkUXBoV0ZJMVRIazFibUZZVW05a1YwbDJaREk1ZVdFeVduTmlNMlI2VERKS01XRlhlR3RNYm14b1lsZDRRV050Vm0xamVUbHZXbGRHYTJONU9YUlpWMngxQ2sxRVowZERhWE5IUVZGUlFtYzNPSGRCVVc5RlMyZDNiMDVIVVRWYWFtTTFUMGRKTlU1WFZtdE5hbWQzV1dwak0wMUVTWGhOZWtFd1QwZEpORnBFVG0wS1RrUkZNVTFFYkd0TlIwMDFXbFJCWWtKbmIzSkNaMFZGUVZsUEwwMUJSVXhDUVRCTlF6Tk9iR0pIV1hSaFJ6bDZaRWRXYTAxRlFVZERhWE5IUVZGUlFncG5OemgzUVZGM1JVMW5kM2RoU0ZJd1kwaE5Oa3g1T1c1aFdGSnZaRmRKZFZreU9YUk1NbWd4V2pKa2NHSnRaRzFaVjA1c1RESjBiR050Tld4aVNFMTBDbGt5T1hSaVdGWjFZVmhTTlUxRVowZERhWE5IUVZGUlFtYzNPSGRCVVRCRlMyZDNiMDVIVVRWYWFtTTFUMGRKTlU1WFZtdE5hbWQzV1dwak0wMUVTWGdLVFhwQk1FOUhTVFJhUkU1dFRrUkZNVTFFYkd0TlIwMDFXbFJCWmtKbmIzSkNaMFZGUVZsUEwwMUJSVTlDUWtWTlJETktiRnB1VFhaaFIxWm9Xa2hOZGdwaVYwWndZbXBCWVVKbmIzSkNaMFZGUVZsUEwwMUJSVkJDUVhkTlEycEZkMDU2UlRCT2VsVXhUV3ByZDB4bldVdExkMWxDUWtGSFJIWjZRVUpGUVZGbkNrUkNOVzlrU0ZKM1kzcHZka3d5WkhCa1IyZ3hXV2sxYW1JeU1IWmhTRlp1V2pKc2RWb3lXbWhaTWxWM1IwRlpTMHQzV1VKQ1FVZEVkbnBCUWtWUlVVc0tSRUZuZVU1VVkzbE5SR013VFhwQ2RFSm5iM0pDWjBWRlFWbFBMMDFCUlZOQ1JqaE5XRmRvTUdSSVFucFBhVGgyV2pKc01HRklWbWxNYlU1MllsTTVid3BrVjJSdVlWYzFibHB0Um1wYVV6bHlXbGhLZFZwWGVIcE1WMDUyWWxjeE1XSnRiREJsVXpoMVdqSnNNR0ZJVm1sTU0yUjJZMjEwYldKSE9UTmplVGxwQ21SWGJITmFRelUxV1ZjeGMxRklTbXhhYmsxMllVZFdhRnBJVFhaaVYwWndZbXBCTkVKbmIzSkNaMFZGUVZsUEwwMUJSVlJDUTI5TlMwUlNhMDlYV1RNS1QxUm9hVTlVVm14YVJFazBUVWRKTTA1NlFYbE5WRTEzVGtSb2FVOUhVWHBhYWxGNFRsUkJOVnBFUW1wUFYxVjNTVkZaUzB0M1dVSkNRVWRFZG5wQlFncEdRVkZVUkVKR00ySXpTbkphYlhoMlpERTVhMkZZVG5kWldGSnFZVVJDYTBKbmIzSkNaMFZGUVZsUEwwMUJSVlpDUmxsTlZrZG9NR1JJUW5wUGFUaDJDbG95YkRCaFNGWnBURzFPZG1KVE9XOWtWMlJ1WVZjMWJscHRSbXBhVXpseVdsaEtkVnBYZUhwTVYwNTJZbGN4TVdKdGJEQmxVemxvV1ROU2NHSXlOWG9LVEROS01XSnVUWFpOYW10NlRYcEJNazFxVFRST2FsVjJXVmhTTUZwWE1YZGtTRTEyVFZSQlYwSm5iM0pDWjBWRlFWbFBMMDFCUlZkQ1FXZE5RbTVDTVFwWmJYaHdXWHBDUjBKbmIzSkNaMFZGUVZsUEwwMUJSVmxDUkdkTlRtNUtiR05IT0RaaFNGWnVXakpzZFZveVdtaFpNbFYyWVRKV2VXSnRWbk5qZVRGcUNtSXlNWFJrVnpWd1pFaHJObU50Vm0xUGJrcHNXbTVOZG1GSFZtaGFTRTEyWWxkR2NHSnFRMEpwWjFsTFMzZFpRa0pCU0ZkbFVVbEZRV2RTT0VKSWIwRUtaVUZDTWtGT01EbE5SM0pIZUhoRmVWbDRhMlZJU214dVRuZExhVk5zTmpRemFubDBMelJsUzJOdlFYWkxaVFpQUVVGQlFtNHlRME5MWWtWQlFVRlJSQXBCUldOM1VsRkpaMDF2VmxVelpIUkNia2RzZUUxTllqTnpUMVJOT0drdmNrOWthWEZoVVRNM2VEQjBlamxPVDFwaEswMURTVkZEZG5KSVVuYzJaa0kxQ21zM2NtOXJUV1UyYVVSMksyNDVjMmxDWldWRlJFazFTR1ZzWXpSME9GQk9abnBCUzBKblozRm9hMnBQVUZGUlJFRjNUbkJCUkVKdFFXcEZRWFZoZW1NS1FtbFVZM1k0U3paTmVVTmllRXByVHpoUk1WZFlOelV5SzA1QldVWTJkVkZvVWtWb1QyeG1URmt4U0hSbU1rRm5WVFZaZW5oWGQzWTNNbTVUUVdwRlFRcHplVWxoWjFSeWQzSndjREJvY1U1UFVVazVhV3MzUTJ3MGNrMXhjSEp3VVRrd2QwdFlSWG8xSzNWeWNWWnRaWFJVYld0a04xbHpUVUZ5UkhweWFVUnZDaTB0TFMwdFJVNUVJRU5GVWxSSlJrbERRVlJGTFMwdExTMEsifX19fQ=="}], "timestampVerificationData":{"rfc3161Timestamps":[{"signedTimestamp":"MIICyjADAgEAMIICwQYJKoZIhvcNAQcCoIICsjCCAq4CAQMxDTALBglghkgBZQMEAgEwgbcGCyqGSIb3DQEJEAEEoIGnBIGkMIGhAgEBBgkrBgEEAYO/MAIwMTANBglghkgBZQMEAgEFAAQg/3UMGV8fwrsfiQEAzqbUSlYc87EUEy8wrOfhaOM5NwACFBwa0r74wyBkjdrXOXL45py8exMrGA8yMDI2MDcxNDEyMDI1MVowAwIBAaAypDAwLjEVMBMGA1UEChMMc2lnc3RvcmUuZGV2MRUwEwYDVQQDEwxzaWdzdG9yZS10c2GgADGCAdwwggHYAgEBMFEwOTEVMBMGA1UEChMMc2lnc3RvcmUuZGV2MSAwHgYDVQQDExdzaWdzdG9yZS10c2Etc2VsZnNpZ25lZAIUOhNULwyQYe68wUMvy4qOiyojiwwwCwYJYIZIAWUDBAIBoIH8MBoGCSqGSIb3DQEJAzENBgsqhkiG9w0BCRABBDAcBgkqhkiG9w0BCQUxDxcNMjYwNzE0MTIwMjUxWjAvBgkqhkiG9w0BCQQxIgQgC5Qp2/4m1NXACJdcq25HSjgek+k/yuqiVNfBceqic8kwgY4GCyqGSIb3DQEJEAIvMX8wfTB7MHkEIIX5J7wHq2LKw7RDVsEO/IGyxog/2nq55thw2dE6zQW3MFUwPaQ7MDkxFTATBgNVBAoTDHNpZ3N0b3JlLmRldjEgMB4GA1UEAxMXc2lnc3RvcmUtdHNhLXNlbGZzaWduZWQCFDoTVC8MkGHuvMFDL8uKjosqI4sMMAoGCCqGSM49BAMCBGgwZgIxALBGy+9QLpTAvIjqpcf3WnWsVHfTrPKGAClJboW2OzfnuQ9wBO4Qi9+XIsmEIuGPdwIxAKbZ0fUnZ4D5AG3RMAvGaeLMf2WAwCkkT+7PEpjMl7uXe2luLTuwr+SqoAxPJxE2KA=="}]}}, "messageSignature":{"messageDigest":{"algorithm":"SHA2_256", "digest":"tr0SClqM2qQrkZGvDA7UNDEBmebN3YXBpRhQGBfHZJs="}, "signature":"MEYCIQCG2a5cjD/Oscu+vZqZ5gmpzKlULhWdxrFeRHM4vTIkJwIhAK+mZA5ZPYPl+frTHWKbNSeLEhXyN7sguB6jL1lALHGL"}}
|
|
|
|
| 1 |
+
{"mediaType":"application/vnd.dev.sigstore.bundle.v0.3+json", "verificationMaterial":{"certificate":{"rawBytes":"MIIHSzCCBtGgAwIBAgIUO6RSJ/EnqV44CIEDdOOD20lhcqIwCgYIKoZIzj0EAwMwNzEVMBMGA1UEChMMc2lnc3RvcmUuZGV2MR4wHAYDVQQDExVzaWdzdG9yZS1pbnRlcm1lZGlhdGUwHhcNMjYwNzIwMTMzMzA2WhcNMjYwNzIwMTM0MzA2WjAAMFkwEwYHKoZIzj0CAQYIKoZIzj0DAQcDQgAE8cQGo4I4BMQ8MJoUyiqSGUWjUiCqFqm2tJO888SUpnLtnkFo1FzBQjCHpomiYGRFWwqO0B19Wr1fIelQoxsFUKOCBfAwggXsMA4GA1UdDwEB/wQEAwIHgDATBgNVHSUEDDAKBggrBgEFBQcDAzAdBgNVHQ4EFgQUv6O4N4tqJqgSLVrYxeRJ/tmXx2UwHwYDVR0jBBgwFoAU39Ppz1YkEZb5qNjpKFWixi4YZD8wawYDVR0RAQH/BGEwX4ZdaHR0cHM6Ly9naXRodWIuY29tL2h1Z2dpbmdmYWNlL2tlcm5lbHMtY29tbXVuaXR5Ly5naXRodWIvd29ya2Zsb3dzL2J1aWxkLnlhbWxAcmVmcy9oZWFkcy9tYWluMDkGCisGAQQBg78wAQEEK2h0dHBzOi8vdG9rZW4uYWN0aW9ucy5naXRodWJ1c2VyY29udGVudC5jb20wHwYKKwYBBAGDvzABAgQRd29ya2Zsb3dfZGlzcGF0Y2gwNgYKKwYBBAGDvzABAwQoMGM1ZmIzM2QyMTA4NGUwNzc1MGEyNDhkYTEyMDcxZmYwMTFiZDcwYTATBgorBgEEAYO/MAEEBAVCdWlsZDArBgorBgEEAYO/MAEFBB1odWdnaW5nZmFjZS9rZXJuZWxzLWNvbW11bml0eTAdBgorBgEEAYO/MAEGBA9yZWZzL2hlYWRzL21haW4wOwYKKwYBBAGDvzABCAQtDCtodHRwczovL3Rva2VuLmFjdGlvbnMuZ2l0aHVidXNlcmNvbnRlbnQuY29tMG0GCisGAQQBg78wAQkEXwxdaHR0cHM6Ly9naXRodWIuY29tL2h1Z2dpbmdmYWNlL2tlcm5lbHMtY29tbXVuaXR5Ly5naXRodWIvd29ya2Zsb3dzL2J1aWxkLnlhbWxAcmVmcy9oZWFkcy9tYWluMDgGCisGAQQBg78wAQoEKgwoMGM1ZmIzM2QyMTA4NGUwNzc1MGEyNDhkYTEyMDcxZmYwMTFiZDcwYTAbBgorBgEEAYO/MAELBA0MC3NlbGYtaG9zdGVkMEAGCisGAQQBg78wAQwEMgwwaHR0cHM6Ly9naXRodWIuY29tL2h1Z2dpbmdmYWNlL2tlcm5lbHMtY29tbXVuaXR5MDgGCisGAQQBg78wAQ0EKgwoMGM1ZmIzM2QyMTA4NGUwNzc1MGEyNDhkYTEyMDcxZmYwMTFiZDcwYTAfBgorBgEEAYO/MAEOBBEMD3JlZnMvaGVhZHMvbWFpbjAaBgorBgEEAYO/MAEPBAwMCjEwNzE0NzU1MjkwLgYKKwYBBAGDvzABEAQgDB5odHRwczovL2dpdGh1Yi5jb20vaHVnZ2luZ2ZhY2UwGAYKKwYBBAGDvzABEQQKDAgyNTcyMDc0MzBtBgorBgEEAYO/MAESBF8MXWh0dHBzOi8vZ2l0aHViLmNvbS9odWdnaW5nZmFjZS9rZXJuZWxzLWNvbW11bml0eS8uZ2l0aHViL3dvcmtmbG93cy9idWlsZC55YW1sQHJlZnMvaGVhZHMvbWFpbjA4BgorBgEEAYO/MAETBCoMKDBjNWZiMzNkMjEwODRlMDc3NTBhMjQ4ZGExMjA3MWZmMDExYmQ3MGEwIQYKKwYBBAGDvzABFAQTDBF3b3JrZmxvd19kaXNwYXRjaDBkBgorBgEEAYO/MAEVBFYMVGh0dHBzOi8vZ2l0aHViLmNvbS9odWdnaW5nZmFjZS9rZXJuZWxzLWNvbW11bml0eS9hY3Rpb25zL3J1bnMvMjk3NDYyMzc1MzUvYXR0ZW1wdHMvMTAWBgorBgEEAYO/MAEWBAgMBnB1YmxpYzBGBgorBgEEAYO/MAEYBDgMNnJlcG86aHVnZ2luZ2ZhY2Uva2VybmVscy1jb21tdW5pdHk6cmVmOnJlZnMvaGVhZHMvbWFpbjCBigYKKwYBBAHWeQIEAgR8BHoAeAB2AN09MGrGxxEyYxkeHJlnNwKiSl643jyt/4eKcoAvKe6OAAABn3+68ykAAAQDAEcwRQIhAKfi0uLcjC+7sKlx6lksgEPD/NhuLlgiGj7ibyikciwIAiBL8jxz6LA5Ggrj5nN1SaUm/ttST2HZx2F8xHAd1hu32zAKBggqhkjOPQQDAwNoADBlAjBowDB5OdvLqQ2GqWrG8F+E7lNV2bbvOsz5keEUywdb1JKbKczCW+ZE5TPgasmM0skCMQC5VBemdmpIL4GaKEgobQkePdAc/wQcmTF743v/+wJ6rtHl1qO6tRVmurEflsSKryQ="}, "tlogEntries":[{"logIndex":"2206355914", "logId":{"keyId":"wNI9atQGlz+VWfO6LRygH4QUfY/8W4RFwiT5i5WRgB0="}, "kindVersion":{"kind":"hashedrekord", "version":"0.0.1"}, "integratedTime":"1784554386", "inclusionPromise":{"signedEntryTimestamp":"MEYCIQDG3iFXAb+fjoVxsKm13pM8c5eEoZl7sIUSR7PCOJQPJwIhAMNN6Gc2ZifBnxZxAFkqOD6fZ9pxz6eIigLvWieBe+12"}, "inclusionProof":{"logIndex":"2084451652", "rootHash":"zbQwyp11BrPe3wggHlyhiEtVAaHleIQPjvxj/sBE5zI=", "treeSize":"2084451654", "hashes":["499kcb89L18pO0S0b4IQogj4NQ5pOzw5FHU004db5jY=", "sd8iJSY7iV+Ikc42LJKo8JVRb2fypg73SOufXNOi0Xo=", "p+vUmE/YgN+ll593XYjUalho2bs1PslkYfN9ocL4t08=", "uandAueMLs9RC3xIj77ePNadPcbKcEeORWqtfLFvAu8=", "SjVrRtIeid4t+H6u4drFM6V4zL5B3xFoI1cbwfFF6kI=", "oz8/zcwvzjAiODVBCpcZm4tSU+3KvdHqIQqxFRSS418=", "LH3MOzxultOUBpT2DD8FsK1ItkTk2kLtqOkxxthdBAE=", "NOqu2omVaTD0z/phbf6aHupVZtPpniJjH7Mohp9dQVg=", "UaVIhKZL9YPl5aIrrdhBcDNZES0kGKGGU+n9C0aUu04=", "S6fX78YW7K/Nsie8dVZY6CYd6Qnvj92xsSI5UOFLMIs=", "phX4Nc7EXYzCDd9JqHuZvmMmWgixvVF/zVGv5sTGpx4=", "V6wr2bDyz/qjnYx99CtUrvwLeHOpXEswquva8IIXU8U=", "GxT6+CEXxY8Ak/besnRW/IK3DhneHQ5X5rIcw8L01DQ=", "Rq/A2aTC9e54ldNjcpsJ26rX/h8JlN6ZHxgw7yEa6C0=", "+/VZ56MsIPxMiyLAodzKXo5TEWdQp36z89qLhpzloAo=", "daxmZaajRpZV+JxHiOYZhJBiSKN5ucqjh2WnGbHhirw=", "DOCeoSMovIvLExkhIvisow9AuNXgeWs4ECkyR6EcqYU="], "checkpoint":{"envelope":"rekor.sigstore.dev - 1193050959916656506\n2084451654\nzbQwyp11BrPe3wggHlyhiEtVAaHleIQPjvxj/sBE5zI=\n\n— rekor.sigstore.dev wNI9ajBFAiAN5Kb1loICwLlR0NDXi67Zk+v9N/w+2dwjAk1gfEA7QwIhAKxiGa8YCzAVUYhlDnLZ/mDkJsC0CNlX2XHsegUBfTIa\n"}}, "canonicalizedBody":"eyJhcGlWZXJzaW9uIjoiMC4wLjEiLCJraW5kIjoiaGFzaGVkcmVrb3JkIiwic3BlYyI6eyJkYXRhIjp7Imhhc2giOnsiYWxnb3JpdGhtIjoic2hhMjU2IiwidmFsdWUiOiI2Y2I2NDk4NTk2YzMwNDFiMWE5MTQzMDI5MzkwNTBmNmUxOGJiYzMzNmFjNDM4YjI3NTA3MGMyYzg1YjJkNzhlIn19LCJzaWduYXR1cmUiOnsiY29udGVudCI6Ik1FWUNJUUNjZ3NYdHNJWWxWRTJ4YVZtcDJkdzJpdjNPZzlDK0MvbEV2NUtyZUxEUDhRSWhBTGxEWUxmaGkrY0FDL2lwUnUxaHMralA2aHdIWU5aR1FTclA5a1lFcHMyVyIsInB1YmxpY0tleSI6eyJjb250ZW50IjoiTFMwdExTMUNSVWRKVGlCRFJWSlVTVVpKUTBGVVJTMHRMUzB0Q2sxSlNVaFRla05EUW5SSFowRjNTVUpCWjBsVlR6WlNVMG92Ulc1eFZqUTBRMGxGUkdSUFQwUXlNR3hvWTNGSmQwTm5XVWxMYjFwSmVtb3dSVUYzVFhjS1RucEZWazFDVFVkQk1WVkZRMmhOVFdNeWJHNWpNMUoyWTIxVmRWcEhWakpOVWpSM1NFRlpSRlpSVVVSRmVGWjZZVmRrZW1SSE9YbGFVekZ3WW01U2JBcGpiVEZzV2tkc2FHUkhWWGRJYUdOT1RXcFpkMDU2U1hkTlZFMTZUWHBCTWxkb1kwNU5hbGwzVG5wSmQwMVVUVEJOZWtFeVYycEJRVTFHYTNkRmQxbElDa3R2V2tsNmFqQkRRVkZaU1V0dldrbDZhakJFUVZGalJGRm5RVVU0WTFGSGJ6UkpORUpOVVRoTlNtOVZlV2x4VTBkVlYycFZhVU54Um5GdE1uUktUemdLT0RoVFZYQnVUSFJ1YTBadk1VWjZRbEZxUTBod2IyMXBXVWRTUmxkM2NVOHdRakU1VjNJeFprbGxiRkZ2ZUhOR1ZVdFBRMEptUVhkbloxaHpUVUUwUndwQk1WVmtSSGRGUWk5M1VVVkJkMGxJWjBSQlZFSm5UbFpJVTFWRlJFUkJTMEpuWjNKQ1owVkdRbEZqUkVGNlFXUkNaMDVXU0ZFMFJVWm5VVlYyTms4MENrNDBkSEZLY1dkVFRGWnlXWGhsVWtvdmRHMVllREpWZDBoM1dVUldVakJxUWtKbmQwWnZRVlV6T1ZCd2VqRlphMFZhWWpWeFRtcHdTMFpYYVhocE5Ga0tXa1E0ZDJGM1dVUldVakJTUVZGSUwwSkhSWGRZTkZwa1lVaFNNR05JVFRaTWVUbHVZVmhTYjJSWFNYVlpNamwwVERKb01Wb3laSEJpYldSdFdWZE9iQXBNTW5Sc1kyMDFiR0pJVFhSWk1qbDBZbGhXZFdGWVVqVk1lVFZ1WVZoU2IyUlhTWFprTWpsNVlUSmFjMkl6WkhwTU1rb3hZVmQ0YTB4dWJHaGlWM2hCQ21OdFZtMWplVGx2V2xkR2EyTjVPWFJaVjJ4MVRVUnJSME5wYzBkQlVWRkNaemM0ZDBGUlJVVkxNbWd3WkVoQ2VrOXBPSFprUnpseVdsYzBkVmxYVGpBS1lWYzVkV041Tlc1aFdGSnZaRmRLTVdNeVZubFpNamwxWkVkV2RXUkROV3BpTWpCM1NIZFpTMHQzV1VKQ1FVZEVkbnBCUWtGblVWSmtNamw1WVRKYWN3cGlNMlJtV2tkc2VtTkhSakJaTW1kM1RtZFpTMHQzV1VKQ1FVZEVkbnBCUWtGM1VXOU5SMDB4V20xSmVrMHlVWGxOVkVFMFRrZFZkMDU2WXpGTlIwVjVDazVFYUd0WlZFVjVUVVJqZUZwdFdYZE5WRVpwV2tSamQxbFVRVlJDWjI5eVFtZEZSVUZaVHk5TlFVVkZRa0ZXUTJSWGJITmFSRUZ5UW1kdmNrSm5SVVVLUVZsUEwwMUJSVVpDUWpGdlpGZGtibUZYTlc1YWJVWnFXbE01Y2xwWVNuVmFWM2g2VEZkT2RtSlhNVEZpYld3d1pWUkJaRUpuYjNKQ1owVkZRVmxQTHdwTlFVVkhRa0U1ZVZwWFducE1NbWhzV1ZkU2Vrd3lNV2hoVnpSM1QzZFpTMHQzV1VKQ1FVZEVkbnBCUWtOQlVYUkVRM1J2WkVoU2QyTjZiM1pNTTFKMkNtRXlWblZNYlVacVpFZHNkbUp1VFhWYU1td3dZVWhXYVdSWVRteGpiVTUyWW01U2JHSnVVWFZaTWpsMFRVY3dSME5wYzBkQlVWRkNaemM0ZDBGUmEwVUtXSGQ0WkdGSVVqQmpTRTAyVEhrNWJtRllVbTlrVjBsMVdUSTVkRXd5YURGYU1tUndZbTFrYlZsWFRteE1NblJzWTIwMWJHSklUWFJaTWpsMFlsaFdkUXBoV0ZJMVRIazFibUZZVW05a1YwbDJaREk1ZVdFeVduTmlNMlI2VERKS01XRlhlR3RNYm14b1lsZDRRV050Vm0xamVUbHZXbGRHYTJONU9YUlpWMngxQ2sxRVowZERhWE5IUVZGUlFtYzNPSGRCVVc5RlMyZDNiMDFIVFRGYWJVbDZUVEpSZVUxVVFUUk9SMVYzVG5wak1VMUhSWGxPUkdocldWUkZlVTFFWTNnS1dtMVpkMDFVUm1sYVJHTjNXVlJCWWtKbmIzSkNaMFZGUVZsUEwwMUJSVXhDUVRCTlF6Tk9iR0pIV1hSaFJ6bDZaRWRXYTAxRlFVZERhWE5IUVZGUlFncG5OemgzUVZGM1JVMW5kM2RoU0ZJd1kwaE5Oa3g1T1c1aFdGSnZaRmRKZFZreU9YUk1NbWd4V2pKa2NHSnRaRzFaVjA1c1RESjBiR050Tld4aVNFMTBDbGt5T1hSaVdGWjFZVmhTTlUxRVowZERhWE5IUVZGUlFtYzNPSGRCVVRCRlMyZDNiMDFIVFRGYWJVbDZUVEpSZVUxVVFUUk9SMVYzVG5wak1VMUhSWGtLVGtSb2ExbFVSWGxOUkdONFdtMVpkMDFVUm1sYVJHTjNXVlJCWmtKbmIzSkNaMFZGUVZsUEwwMUJSVTlDUWtWTlJETktiRnB1VFhaaFIxWm9Xa2hOZGdwaVYwWndZbXBCWVVKbmIzSkNaMFZGUVZsUEwwMUJSVkJDUVhkTlEycEZkMDU2UlRCT2VsVXhUV3ByZDB4bldVdExkMWxDUWtGSFJIWjZRVUpGUVZGbkNrUkNOVzlrU0ZKM1kzcHZka3d5WkhCa1IyZ3hXV2sxYW1JeU1IWmhTRlp1V2pKc2RWb3lXbWhaTWxWM1IwRlpTMHQzV1VKQ1FVZEVkbnBCUWtWUlVVc0tSRUZuZVU1VVkzbE5SR013VFhwQ2RFSm5iM0pDWjBWRlFWbFBMMDFCUlZOQ1JqaE5XRmRvTUdSSVFucFBhVGgyV2pKc01HRklWbWxNYlU1MllsTTVid3BrVjJSdVlWYzFibHB0Um1wYVV6bHlXbGhLZFZwWGVIcE1WMDUyWWxjeE1XSnRiREJsVXpoMVdqSnNNR0ZJVm1sTU0yUjJZMjEwYldKSE9UTmplVGxwQ21SWGJITmFRelUxV1ZjeGMxRklTbXhhYmsxMllVZFdhRnBJVFhaaVYwWndZbXBCTkVKbmIzSkNaMFZGUVZsUEwwMUJSVlJDUTI5TlMwUkNhazVYV21rS1RYcE9hMDFxUlhkUFJGSnNUVVJqTTA1VVFtaE5hbEUwV2tkRmVFMXFRVE5OVjFwdFRVUkZlRmx0VVROTlIwVjNTVkZaUzB0M1dVSkNRVWRFZG5wQlFncEdRVkZVUkVKR00ySXpTbkphYlhoMlpERTVhMkZZVG5kWldGSnFZVVJDYTBKbmIzSkNaMFZGUVZsUEwwMUJSVlpDUmxsTlZrZG9NR1JJUW5wUGFUaDJDbG95YkRCaFNGWnBURzFPZG1KVE9XOWtWMlJ1WVZjMWJscHRSbXBhVXpseVdsaEtkVnBYZUhwTVYwNTJZbGN4TVdKdGJEQmxVemxvV1ROU2NHSXlOWG9LVEROS01XSnVUWFpOYW1zelRrUlplVTE2WXpGTmVsVjJXVmhTTUZwWE1YZGtTRTEyVFZSQlYwSm5iM0pDWjBWRlFWbFBMMDFCUlZkQ1FXZE5RbTVDTVFwWmJYaHdXWHBDUjBKbmIzSkNaMFZGUVZsUEwwMUJSVmxDUkdkTlRtNUtiR05IT0RaaFNGWnVXakpzZFZveVdtaFpNbFYyWVRKV2VXSnRWbk5qZVRGcUNtSXlNWFJrVnpWd1pFaHJObU50Vm0xUGJrcHNXbTVOZG1GSFZtaGFTRTEyWWxkR2NHSnFRMEpwWjFsTFMzZFpRa0pCU0ZkbFVVbEZRV2RTT0VKSWIwRUtaVUZDTWtGT01EbE5SM0pIZUhoRmVWbDRhMlZJU214dVRuZExhVk5zTmpRemFubDBMelJsUzJOdlFYWkxaVFpQUVVGQlFtNHpLelk0ZVd0QlFVRlJSQXBCUldOM1VsRkphRUZMWm1rd2RVeGpha01yTjNOTGJIZzJiR3R6WjBWUVJDOU9hSFZNYkdkcFIybzNhV0o1YVd0amFYZEpRV2xDVERocWVIbzJURUUxQ2tkbmNtbzFiazR4VTJGVmJTOTBkRk5VTWtoYWVESkdPSGhJUVdReGFIVXpNbnBCUzBKblozRm9hMnBQVUZGUlJFRjNUbTlCUkVKc1FXcENiM2RFUWpVS1QyUjJUSEZSTWtkeFYzSkhPRVlyUlRkc1RsWXlZbUoyVDNONk5XdGxSVlY1ZDJSaU1VcExZa3RqZWtOWEsxcEZOVlJRWjJGemJVMHdjMnREVFZGRE5RcFdRbVZ0Wkcxd1NVdzBSMkZMUldkdllsRnJaVkJrUVdNdmQxRmpiVlJHTnpRemRpOHJkMG8yY25SSWJERnhUelowVWxadGRYSkZabXh6VTB0eWVWRTlDaTB0TFMwdFJVNUVJRU5GVWxSSlJrbERRVlJGTFMwdExTMEsifX19fQ=="}], "timestampVerificationData":{"rfc3161Timestamps":[{"signedTimestamp":"MIICyzADAgEAMIICwgYJKoZIhvcNAQcCoIICszCCAq8CAQMxDTALBglghkgBZQMEAgEwgbgGCyqGSIb3DQEJEAEEoIGoBIGlMIGiAgEBBgkrBgEEAYO/MAIwMTANBglghkgBZQMEAgEFAAQgvw1pUaMIhxjbKLh6+5XKQorzGNdAiR35CAi3ahtL/wICFQDUaY9j5Oats6jvX/DfaOYqFh0byxgPMjAyNjA3MjAxMzMzMDZaMAMCAQGgMqQwMC4xFTATBgNVBAoTDHNpZ3N0b3JlLmRldjEVMBMGA1UEAxMMc2lnc3RvcmUtdHNhoAAxggHcMIIB2AIBATBRMDkxFTATBgNVBAoTDHNpZ3N0b3JlLmRldjEgMB4GA1UEAxMXc2lnc3RvcmUtdHNhLXNlbGZzaWduZWQCFDoTVC8MkGHuvMFDL8uKjosqI4sMMAsGCWCGSAFlAwQCAaCB/DAaBgkqhkiG9w0BCQMxDQYLKoZIhvcNAQkQAQQwHAYJKoZIhvcNAQkFMQ8XDTI2MDcyMDEzMzMwNlowLwYJKoZIhvcNAQkEMSIEIPCocok9oFF41l0u0wS8It++csm/ItVfTR/zQ6TitjHTMIGOBgsqhkiG9w0BCRACLzF/MH0wezB5BCCF+Se8B6tiysO0Q1bBDvyBssaIP9p6uebYcNnROs0FtzBVMD2kOzA5MRUwEwYDVQQKEwxzaWdzdG9yZS5kZXYxIDAeBgNVBAMTF3NpZ3N0b3JlLXRzYS1zZWxmc2lnbmVkAhQ6E1QvDJBh7rzBQy/Lio6LKiOLDDAKBggqhkjOPQQDAgRoMGYCMQDUU5At837LJul65e/JIa4/4I1tVXnMU7HV2Y1f3HuM4ddjXQZSbWn8if7ZaoDFenoCMQCEPz2BqJsUGrxCWijUtVD8SJx79reROK3HNxmShRIsaA9ahJiE6V2W137HRpPWsR8="}]}}, "messageSignature":{"messageDigest":{"algorithm":"SHA2_256", "digest":"bLZJhZbDBBsakUMCk5BQ9uGLvDNqxDiydQcMLIWy144="}, "signature":"MEYCIQCcgsXtsIYlVE2xaVmp2dw2iv3Og9C+C/lEv5KreLDP8QIhALlDYLfhi+cAC/ipRu1hs+jP6hwHYNZGQSrP9kYEps2W"}}
|
build/torch-rocm/qwen2vl_mrope.py
CHANGED
|
@@ -2,6 +2,8 @@ import torch
|
|
| 2 |
import triton
|
| 3 |
import triton.language as tl
|
| 4 |
|
|
|
|
|
|
|
| 5 |
|
| 6 |
@triton.jit
|
| 7 |
def _triton_qwen2vl_mrope(
|
|
@@ -128,24 +130,25 @@ def qwen2vl_mrope_forward(q, k, cos, sin, mrope_section):
|
|
| 128 |
cos = cos.contiguous()
|
| 129 |
sin = sin.contiguous()
|
| 130 |
|
| 131 |
-
|
| 132 |
-
|
| 133 |
-
|
| 134 |
-
|
| 135 |
-
|
| 136 |
-
|
| 137 |
-
|
| 138 |
-
|
| 139 |
-
|
| 140 |
-
|
| 141 |
-
|
| 142 |
-
|
| 143 |
-
|
| 144 |
-
|
| 145 |
-
|
| 146 |
-
|
| 147 |
-
|
| 148 |
-
|
|
|
|
| 149 |
return q.transpose(1, 2), k.transpose(1, 2), cos, sin
|
| 150 |
|
| 151 |
|
|
@@ -166,25 +169,26 @@ def qwen2vl_mrope_backward(dq, dk, cos, sin, mrope_section):
|
|
| 166 |
dq = dq.contiguous()
|
| 167 |
dk = dk.contiguous()
|
| 168 |
|
| 169 |
-
|
| 170 |
-
|
| 171 |
-
|
| 172 |
-
|
| 173 |
-
|
| 174 |
-
|
| 175 |
-
|
| 176 |
-
|
| 177 |
-
|
| 178 |
-
|
| 179 |
-
|
| 180 |
-
|
| 181 |
-
|
| 182 |
-
|
| 183 |
-
|
| 184 |
-
|
| 185 |
-
|
| 186 |
-
|
| 187 |
-
|
|
|
|
| 188 |
return dq.transpose(1, 2), dk.transpose(1, 2)
|
| 189 |
|
| 190 |
|
|
|
|
| 2 |
import triton
|
| 3 |
import triton.language as tl
|
| 4 |
|
| 5 |
+
from .utils import device_context
|
| 6 |
+
|
| 7 |
|
| 8 |
@triton.jit
|
| 9 |
def _triton_qwen2vl_mrope(
|
|
|
|
| 130 |
cos = cos.contiguous()
|
| 131 |
sin = sin.contiguous()
|
| 132 |
|
| 133 |
+
with device_context(q.device):
|
| 134 |
+
_triton_qwen2vl_mrope[(n_row,)](
|
| 135 |
+
q,
|
| 136 |
+
k,
|
| 137 |
+
cos,
|
| 138 |
+
sin,
|
| 139 |
+
seq_len,
|
| 140 |
+
batch_size,
|
| 141 |
+
n_q_head,
|
| 142 |
+
n_kv_head,
|
| 143 |
+
head_dim,
|
| 144 |
+
pad_n_q_head,
|
| 145 |
+
pad_n_kv_head,
|
| 146 |
+
pad_hd,
|
| 147 |
+
mrope_section[0],
|
| 148 |
+
mrope_section[1],
|
| 149 |
+
BLOCK_SIZE=BLOCK_SIZE,
|
| 150 |
+
BACKWARD_PASS=False,
|
| 151 |
+
)
|
| 152 |
return q.transpose(1, 2), k.transpose(1, 2), cos, sin
|
| 153 |
|
| 154 |
|
|
|
|
| 169 |
dq = dq.contiguous()
|
| 170 |
dk = dk.contiguous()
|
| 171 |
|
| 172 |
+
with device_context(dq.device):
|
| 173 |
+
# backward is similar to forward except swapping few ops
|
| 174 |
+
_triton_qwen2vl_mrope[(n_row,)](
|
| 175 |
+
dq,
|
| 176 |
+
dk,
|
| 177 |
+
cos,
|
| 178 |
+
sin,
|
| 179 |
+
seq_len,
|
| 180 |
+
batch_size,
|
| 181 |
+
n_q_head,
|
| 182 |
+
n_kv_head,
|
| 183 |
+
head_dim,
|
| 184 |
+
pad_n_q_head,
|
| 185 |
+
pad_n_kv_head,
|
| 186 |
+
pad_hd,
|
| 187 |
+
mrope_section[0],
|
| 188 |
+
mrope_section[1],
|
| 189 |
+
BLOCK_SIZE=BLOCK_SIZE,
|
| 190 |
+
BACKWARD_PASS=True,
|
| 191 |
+
)
|
| 192 |
return dq.transpose(1, 2), dk.transpose(1, 2)
|
| 193 |
|
| 194 |
|
build/torch-rocm/rms_norm.py
CHANGED
|
@@ -24,6 +24,8 @@ from .utils import get_npu_core_count
|
|
| 24 |
from .utils import set_large_grf_mode
|
| 25 |
from .utils import torch_to_triton_dtype
|
| 26 |
from .utils import is_npu_available
|
|
|
|
|
|
|
| 27 |
|
| 28 |
if compare_version("triton", operator.ge, "3.0.0") and not is_npu_available():
|
| 29 |
try:
|
|
@@ -438,47 +440,49 @@ def rms_norm_forward(X, W, eps, offset, casting_mode, row_mode):
|
|
| 438 |
kernel_args = {}
|
| 439 |
if X.device.type == "xpu":
|
| 440 |
set_large_grf_mode(kernel_args)
|
| 441 |
-
|
| 442 |
-
|
| 443 |
-
|
| 444 |
-
|
| 445 |
-
|
| 446 |
-
|
| 447 |
-
|
| 448 |
-
|
| 449 |
-
|
| 450 |
-
|
| 451 |
-
|
| 452 |
-
|
| 453 |
-
|
| 454 |
-
|
| 455 |
-
|
| 456 |
-
|
| 457 |
-
|
| 458 |
-
|
| 459 |
-
|
| 460 |
-
|
| 461 |
-
|
| 462 |
-
|
| 463 |
-
|
| 464 |
-
|
| 465 |
-
|
| 466 |
-
|
| 467 |
-
|
| 468 |
-
|
| 469 |
-
|
| 470 |
-
|
| 471 |
-
|
| 472 |
-
|
| 473 |
-
|
| 474 |
-
|
| 475 |
-
|
| 476 |
-
|
| 477 |
-
|
| 478 |
-
|
| 479 |
-
|
| 480 |
-
|
| 481 |
-
|
|
|
|
|
|
|
| 482 |
return Y.view(*shape), X, RSTD, BLOCK_SIZE, num_warps, casting_mode
|
| 483 |
|
| 484 |
|
|
@@ -519,57 +523,58 @@ def rms_norm_backward(dY, X, W, RSTD, offset, casting_mode, BLOCK_SIZE, num_warp
|
|
| 519 |
if X.device.type == "xpu":
|
| 520 |
set_large_grf_mode(kernel_args)
|
| 521 |
|
| 522 |
-
|
| 523 |
-
|
| 524 |
-
|
| 525 |
-
|
| 526 |
-
|
| 527 |
-
|
| 528 |
-
|
| 529 |
-
|
| 530 |
-
|
| 531 |
-
|
| 532 |
-
|
| 533 |
-
|
| 534 |
-
|
| 535 |
-
|
| 536 |
-
|
| 537 |
-
|
| 538 |
-
|
| 539 |
-
|
| 540 |
-
|
| 541 |
-
|
| 542 |
-
|
| 543 |
-
|
| 544 |
-
|
| 545 |
-
|
| 546 |
-
|
| 547 |
-
|
| 548 |
-
|
| 549 |
-
|
| 550 |
-
|
| 551 |
-
|
| 552 |
-
|
| 553 |
-
|
| 554 |
-
|
| 555 |
-
|
| 556 |
-
|
| 557 |
-
|
| 558 |
-
|
| 559 |
-
|
| 560 |
-
|
| 561 |
-
|
| 562 |
-
|
| 563 |
-
|
| 564 |
-
|
| 565 |
-
|
| 566 |
-
|
| 567 |
-
|
| 568 |
-
|
| 569 |
-
|
| 570 |
-
|
| 571 |
-
|
| 572 |
-
|
|
|
|
| 573 |
dX = dX.view(*shape)
|
| 574 |
|
| 575 |
if elementwise_affine:
|
|
|
|
| 24 |
from .utils import set_large_grf_mode
|
| 25 |
from .utils import torch_to_triton_dtype
|
| 26 |
from .utils import is_npu_available
|
| 27 |
+
from .utils import device_context
|
| 28 |
+
|
| 29 |
|
| 30 |
if compare_version("triton", operator.ge, "3.0.0") and not is_npu_available():
|
| 31 |
try:
|
|
|
|
| 440 |
kernel_args = {}
|
| 441 |
if X.device.type == "xpu":
|
| 442 |
set_large_grf_mode(kernel_args)
|
| 443 |
+
|
| 444 |
+
with device_context(X.device):
|
| 445 |
+
if BLOCK_SIZE > 256 or n_rows < 4096 * 8 or row_mode:
|
| 446 |
+
_rms_norm_forward_kernel[(n_rows,)](
|
| 447 |
+
Y,
|
| 448 |
+
Y.stride(0),
|
| 449 |
+
X,
|
| 450 |
+
X.stride(0),
|
| 451 |
+
W,
|
| 452 |
+
W.stride(0) if elementwise_affine else 0,
|
| 453 |
+
RSTD,
|
| 454 |
+
RSTD.stride(0),
|
| 455 |
+
n_cols,
|
| 456 |
+
eps,
|
| 457 |
+
offset,
|
| 458 |
+
casting_mode,
|
| 459 |
+
elementwise_affine=elementwise_affine,
|
| 460 |
+
BLOCK_SIZE=BLOCK_SIZE,
|
| 461 |
+
num_warps=num_warps,
|
| 462 |
+
**kernel_args, # XPU-specific optimization
|
| 463 |
+
)
|
| 464 |
+
else:
|
| 465 |
+
BLOCK_ROW = 16
|
| 466 |
+
kernel_args["BLOCK_ROW"] = BLOCK_ROW
|
| 467 |
+
_block_rms_norm_forward_kernel[(triton.cdiv(n_rows, BLOCK_ROW),)](
|
| 468 |
+
Y,
|
| 469 |
+
Y.stride(0),
|
| 470 |
+
X,
|
| 471 |
+
X.stride(0),
|
| 472 |
+
W,
|
| 473 |
+
W.stride(0) if elementwise_affine else 0,
|
| 474 |
+
RSTD,
|
| 475 |
+
RSTD.stride(0),
|
| 476 |
+
n_rows,
|
| 477 |
+
n_cols,
|
| 478 |
+
eps,
|
| 479 |
+
offset,
|
| 480 |
+
casting_mode,
|
| 481 |
+
elementwise_affine=elementwise_affine,
|
| 482 |
+
BLOCK_SIZE=BLOCK_SIZE,
|
| 483 |
+
num_warps=num_warps,
|
| 484 |
+
**kernel_args, # XPU-specific optimization
|
| 485 |
+
)
|
| 486 |
return Y.view(*shape), X, RSTD, BLOCK_SIZE, num_warps, casting_mode
|
| 487 |
|
| 488 |
|
|
|
|
| 523 |
if X.device.type == "xpu":
|
| 524 |
set_large_grf_mode(kernel_args)
|
| 525 |
|
| 526 |
+
with device_context(X.device):
|
| 527 |
+
if BLOCK_SIZE > 256 or n_rows < 4096 * 8 or row_mode:
|
| 528 |
+
_rms_norm_backward_kernel[grid](
|
| 529 |
+
dY,
|
| 530 |
+
dY.stride(0),
|
| 531 |
+
dX,
|
| 532 |
+
dX.stride(0),
|
| 533 |
+
X,
|
| 534 |
+
X.stride(0),
|
| 535 |
+
torch_to_triton_dtype[X.dtype],
|
| 536 |
+
W,
|
| 537 |
+
W.stride(0) if elementwise_affine else 0,
|
| 538 |
+
RSTD,
|
| 539 |
+
RSTD.stride(0),
|
| 540 |
+
_dW,
|
| 541 |
+
_dW.stride(0) if elementwise_affine else 0,
|
| 542 |
+
n_rows,
|
| 543 |
+
n_cols,
|
| 544 |
+
offset,
|
| 545 |
+
rows_per_program,
|
| 546 |
+
casting_mode,
|
| 547 |
+
elementwise_affine=elementwise_affine,
|
| 548 |
+
BLOCK_SIZE=BLOCK_SIZE,
|
| 549 |
+
num_warps=num_warps,
|
| 550 |
+
**kernel_args, # XPU-specific optimization
|
| 551 |
+
)
|
| 552 |
+
else:
|
| 553 |
+
BLOCK_ROW = 16
|
| 554 |
+
kernel_args["BLOCK_ROW"] = BLOCK_ROW
|
| 555 |
+
_block_rms_norm_backward_kernel[grid](
|
| 556 |
+
dY,
|
| 557 |
+
dY.stride(0),
|
| 558 |
+
dX,
|
| 559 |
+
dX.stride(0),
|
| 560 |
+
X,
|
| 561 |
+
X.stride(0),
|
| 562 |
+
torch_to_triton_dtype[X.dtype],
|
| 563 |
+
W,
|
| 564 |
+
W.stride(0) if elementwise_affine else 0,
|
| 565 |
+
RSTD,
|
| 566 |
+
RSTD.stride(0),
|
| 567 |
+
_dW,
|
| 568 |
+
_dW.stride(0) if elementwise_affine else 0,
|
| 569 |
+
n_rows,
|
| 570 |
+
n_cols,
|
| 571 |
+
offset,
|
| 572 |
+
casting_mode,
|
| 573 |
+
elementwise_affine=elementwise_affine,
|
| 574 |
+
BLOCK_SIZE=BLOCK_SIZE,
|
| 575 |
+
num_warps=num_warps,
|
| 576 |
+
**kernel_args, # XPU-specific optimization
|
| 577 |
+
)
|
| 578 |
dX = dX.view(*shape)
|
| 579 |
|
| 580 |
if elementwise_affine:
|
build/torch-rocm/rope.py
CHANGED
|
@@ -2,6 +2,8 @@ import torch
|
|
| 2 |
import triton
|
| 3 |
import triton.language as tl
|
| 4 |
|
|
|
|
|
|
|
| 5 |
|
| 6 |
@triton.jit
|
| 7 |
def _triton_rope(
|
|
@@ -134,27 +136,28 @@ def rope_forward(q, k, cos, sin):
|
|
| 134 |
sin = sin.contiguous()
|
| 135 |
cos_batch_size = cos.shape[0]
|
| 136 |
|
| 137 |
-
|
| 138 |
-
|
| 139 |
-
|
| 140 |
-
|
| 141 |
-
|
| 142 |
-
|
| 143 |
-
|
| 144 |
-
|
| 145 |
-
|
| 146 |
-
|
| 147 |
-
|
| 148 |
-
|
| 149 |
-
|
| 150 |
-
|
| 151 |
-
|
| 152 |
-
|
| 153 |
-
|
| 154 |
-
|
| 155 |
-
|
| 156 |
-
|
| 157 |
-
|
|
|
|
| 158 |
return q.transpose(1, 2), k.transpose(1, 2), cos, sin
|
| 159 |
|
| 160 |
|
|
@@ -176,28 +179,29 @@ def rope_backward(dq, dk, cos, sin):
|
|
| 176 |
dq = dq.contiguous()
|
| 177 |
dk = dk.contiguous()
|
| 178 |
|
| 179 |
-
|
| 180 |
-
|
| 181 |
-
|
| 182 |
-
|
| 183 |
-
|
| 184 |
-
|
| 185 |
-
|
| 186 |
-
|
| 187 |
-
|
| 188 |
-
|
| 189 |
-
|
| 190 |
-
|
| 191 |
-
|
| 192 |
-
|
| 193 |
-
|
| 194 |
-
|
| 195 |
-
|
| 196 |
-
|
| 197 |
-
|
| 198 |
-
|
| 199 |
-
|
| 200 |
-
|
|
|
|
| 201 |
return dq.transpose(1, 2), dk.transpose(1, 2)
|
| 202 |
|
| 203 |
|
|
|
|
| 2 |
import triton
|
| 3 |
import triton.language as tl
|
| 4 |
|
| 5 |
+
from .utils import device_context
|
| 6 |
+
|
| 7 |
|
| 8 |
@triton.jit
|
| 9 |
def _triton_rope(
|
|
|
|
| 136 |
sin = sin.contiguous()
|
| 137 |
cos_batch_size = cos.shape[0]
|
| 138 |
|
| 139 |
+
with device_context(q.device):
|
| 140 |
+
_triton_rope[(n_row,)](
|
| 141 |
+
q,
|
| 142 |
+
q.stride(1),
|
| 143 |
+
k,
|
| 144 |
+
k.stride(1),
|
| 145 |
+
cos,
|
| 146 |
+
cos.stride(-2),
|
| 147 |
+
sin,
|
| 148 |
+
sin.stride(-2),
|
| 149 |
+
seq_len,
|
| 150 |
+
batch_size,
|
| 151 |
+
cos_batch_size,
|
| 152 |
+
n_q_head,
|
| 153 |
+
n_kv_head,
|
| 154 |
+
head_dim,
|
| 155 |
+
pad_n_q_head,
|
| 156 |
+
pad_n_kv_head,
|
| 157 |
+
pad_hd,
|
| 158 |
+
BLOCK_SIZE=BLOCK_SIZE,
|
| 159 |
+
BACKWARD_PASS=False,
|
| 160 |
+
)
|
| 161 |
return q.transpose(1, 2), k.transpose(1, 2), cos, sin
|
| 162 |
|
| 163 |
|
|
|
|
| 179 |
dq = dq.contiguous()
|
| 180 |
dk = dk.contiguous()
|
| 181 |
|
| 182 |
+
with device_context(dq.device):
|
| 183 |
+
# backward is similar to forward except swapping few ops
|
| 184 |
+
_triton_rope[(n_row,)](
|
| 185 |
+
dq,
|
| 186 |
+
dq.stride(1),
|
| 187 |
+
dk,
|
| 188 |
+
dk.stride(1),
|
| 189 |
+
cos,
|
| 190 |
+
cos.stride(-2),
|
| 191 |
+
sin,
|
| 192 |
+
sin.stride(-2),
|
| 193 |
+
seq_len,
|
| 194 |
+
batch_size,
|
| 195 |
+
cos_batch_size,
|
| 196 |
+
n_q_head,
|
| 197 |
+
n_kv_head,
|
| 198 |
+
head_dim,
|
| 199 |
+
pad_n_q_head,
|
| 200 |
+
pad_n_kv_head,
|
| 201 |
+
pad_hd,
|
| 202 |
+
BLOCK_SIZE=BLOCK_SIZE,
|
| 203 |
+
BACKWARD_PASS=True,
|
| 204 |
+
)
|
| 205 |
return dq.transpose(1, 2), dk.transpose(1, 2)
|
| 206 |
|
| 207 |
|
build/torch-rocm/swiglu.py
CHANGED
|
@@ -4,6 +4,7 @@ import triton.language as tl
|
|
| 4 |
|
| 5 |
from .utils import calculate_settings
|
| 6 |
from .utils import ensure_contiguous
|
|
|
|
| 7 |
|
| 8 |
|
| 9 |
@triton.jit
|
|
@@ -73,16 +74,17 @@ def swiglu_forward(a, b, gate_multiplier: float = 1.0):
|
|
| 73 |
|
| 74 |
BLOCK_SIZE, num_warps = calculate_settings(n_cols)
|
| 75 |
|
| 76 |
-
|
| 77 |
-
|
| 78 |
-
|
| 79 |
-
|
| 80 |
-
|
| 81 |
-
|
| 82 |
-
|
| 83 |
-
|
| 84 |
-
|
| 85 |
-
|
|
|
|
| 86 |
return a, b, c.view(*ori_shape)
|
| 87 |
|
| 88 |
|
|
@@ -94,16 +96,17 @@ def swiglu_backward(a, b, dc, gate_multiplier: float = 1.0):
|
|
| 94 |
|
| 95 |
BLOCK_SIZE, num_warps = calculate_settings(n_cols)
|
| 96 |
|
| 97 |
-
|
| 98 |
-
|
| 99 |
-
|
| 100 |
-
|
| 101 |
-
|
| 102 |
-
|
| 103 |
-
|
| 104 |
-
|
| 105 |
-
|
| 106 |
-
|
|
|
|
| 107 |
return a.view(*ori_shape), b.view(*ori_shape)
|
| 108 |
|
| 109 |
|
|
|
|
| 4 |
|
| 5 |
from .utils import calculate_settings
|
| 6 |
from .utils import ensure_contiguous
|
| 7 |
+
from .utils import device_context
|
| 8 |
|
| 9 |
|
| 10 |
@triton.jit
|
|
|
|
| 74 |
|
| 75 |
BLOCK_SIZE, num_warps = calculate_settings(n_cols)
|
| 76 |
|
| 77 |
+
with device_context(a.device):
|
| 78 |
+
_swiglu_forward_kernel[(n_rows,)](
|
| 79 |
+
a,
|
| 80 |
+
b,
|
| 81 |
+
c,
|
| 82 |
+
c.stride(-2),
|
| 83 |
+
float(gate_multiplier),
|
| 84 |
+
n_cols=n_cols,
|
| 85 |
+
BLOCK_SIZE=BLOCK_SIZE,
|
| 86 |
+
num_warps=num_warps,
|
| 87 |
+
)
|
| 88 |
return a, b, c.view(*ori_shape)
|
| 89 |
|
| 90 |
|
|
|
|
| 96 |
|
| 97 |
BLOCK_SIZE, num_warps = calculate_settings(n_cols)
|
| 98 |
|
| 99 |
+
with device_context(a.device):
|
| 100 |
+
_swiglu_backward_kernel[(n_rows,)](
|
| 101 |
+
dc,
|
| 102 |
+
a,
|
| 103 |
+
b,
|
| 104 |
+
dc.stride(-2),
|
| 105 |
+
float(gate_multiplier),
|
| 106 |
+
n_cols=n_cols,
|
| 107 |
+
BLOCK_SIZE=BLOCK_SIZE,
|
| 108 |
+
num_warps=num_warps,
|
| 109 |
+
)
|
| 110 |
return a.view(*ori_shape), b.view(*ori_shape)
|
| 111 |
|
| 112 |
|
build/torch-rocm/tvd.py
CHANGED
|
@@ -6,6 +6,7 @@ import triton
|
|
| 6 |
import triton.language as tl
|
| 7 |
|
| 8 |
from .utils import ensure_contiguous
|
|
|
|
| 9 |
|
| 10 |
MAX_FUSED_SIZE = 65536 // 4
|
| 11 |
|
|
@@ -124,24 +125,25 @@ def tv_distance_forward_triton(p, q, shift_labels, reduction, ignore_index, has_
|
|
| 124 |
else:
|
| 125 |
scale = 1.0
|
| 126 |
|
| 127 |
-
|
| 128 |
-
|
| 129 |
-
|
| 130 |
-
|
| 131 |
-
|
| 132 |
-
|
| 133 |
-
|
| 134 |
-
|
| 135 |
-
|
| 136 |
-
|
| 137 |
-
|
| 138 |
-
|
| 139 |
-
|
| 140 |
-
|
| 141 |
-
|
| 142 |
-
|
| 143 |
-
|
| 144 |
-
|
|
|
|
| 145 |
|
| 146 |
# Loss and gradients are already scaled inside the kernel — no separate division needed
|
| 147 |
if reduction in (_REDUCTION_MODE_BATCHMEAN.value, _REDUCTION_MODE_MEAN.value):
|
|
|
|
| 6 |
import triton.language as tl
|
| 7 |
|
| 8 |
from .utils import ensure_contiguous
|
| 9 |
+
from .utils import device_context
|
| 10 |
|
| 11 |
MAX_FUSED_SIZE = 65536 // 4
|
| 12 |
|
|
|
|
| 125 |
else:
|
| 126 |
scale = 1.0
|
| 127 |
|
| 128 |
+
with device_context(p.device):
|
| 129 |
+
_tv_distance_kernel[grid](
|
| 130 |
+
p,
|
| 131 |
+
p.stride(0),
|
| 132 |
+
q,
|
| 133 |
+
q.stride(0),
|
| 134 |
+
output_tensor,
|
| 135 |
+
output_tensor.stride(0),
|
| 136 |
+
grads,
|
| 137 |
+
grads.stride(0),
|
| 138 |
+
shift_labels if has_label else torch.empty(1, device=p.device),
|
| 139 |
+
ignore_index,
|
| 140 |
+
V,
|
| 141 |
+
scale,
|
| 142 |
+
BLOCK_SIZE=BLOCK_SIZE,
|
| 143 |
+
HAS_LABEL=has_label,
|
| 144 |
+
num_warps=num_warps,
|
| 145 |
+
reduction=reduction,
|
| 146 |
+
)
|
| 147 |
|
| 148 |
# Loss and gradients are already scaled inside the kernel — no separate division needed
|
| 149 |
if reduction in (_REDUCTION_MODE_BATCHMEAN.value, _REDUCTION_MODE_MEAN.value):
|
build/torch-rocm/utils.py
CHANGED
|
@@ -21,6 +21,7 @@ import triton
|
|
| 21 |
import triton.language as tl
|
| 22 |
|
| 23 |
from packaging.version import Version
|
|
|
|
| 24 |
|
| 25 |
|
| 26 |
def is_npu_available() -> bool:
|
|
@@ -174,3 +175,13 @@ def set_large_grf_mode(kernel_args: dict):
|
|
| 174 |
else:
|
| 175 |
# API was changed in https://github.com/intel/intel-xpu-backend-for-triton/pull/5430
|
| 176 |
kernel_args["grf_mode"] = "large"
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 21 |
import triton.language as tl
|
| 22 |
|
| 23 |
from packaging.version import Version
|
| 24 |
+
from contextlib import contextmanager
|
| 25 |
|
| 26 |
|
| 27 |
def is_npu_available() -> bool:
|
|
|
|
| 175 |
else:
|
| 176 |
# API was changed in https://github.com/intel/intel-xpu-backend-for-triton/pull/5430
|
| 177 |
kernel_args["grf_mode"] = "large"
|
| 178 |
+
|
| 179 |
+
@contextmanager
|
| 180 |
+
def device_context(device: torch.device):
|
| 181 |
+
"""Context manager that sets the active device for any backend (cuda, xpu, etc.)."""
|
| 182 |
+
backend = getattr(torch, device.type, None)
|
| 183 |
+
if backend is not None and hasattr(backend, "device"):
|
| 184 |
+
with backend.device(device):
|
| 185 |
+
yield
|
| 186 |
+
else:
|
| 187 |
+
yield
|