kernels-bot commited on
Commit
5c41ca4
·
verified ·
1 Parent(s): 7c16513

Uploaded using `kernel-builder`.

Browse files
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 = "4d9f798"
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
- # Here we use a trick to store X_ptr gradient in X_ptr so we can save memory
376
- liger_cross_entropy_kernel[(n_rows,)](
377
- X_ptr=_input,
378
- X_stride=_input.stride(-2),
379
- Y_ptr=target,
380
- Y_stride=target.stride(-1), # always 1
381
- weight_ptr=weight, # dummy if None
382
- loss_ptr=loss_1d,
383
- z_loss_ptr=z_loss_1d,
384
- loss_stride=loss_1d.stride(-1), # always 1
385
- token_accuracy_ptr=token_accuracy_1d,
386
- token_accuracy_stride=token_accuracy_1d.stride(-1)
387
- if return_token_accuracy
388
- else 0, # always 1 if accuracy is enabled
389
- predicted_tokens_ptr=predicted_tokens_1d,
390
- predicted_tokens_stride=predicted_tokens_1d.stride(-1)
391
- if return_predicted_tokens
392
- else 0, # always 1 if predicted tokens is enabled
393
- n_cols=V,
394
- n_non_ignore=n_non_ignore,
395
- sum_non_ignore_weight=sum_non_ignore_weight,
396
- ignore_index=ignore_index,
397
- weight_sum=weight_sum,
398
- lse_square_scale=lse_square_scale,
399
- label_smoothing=label_smoothing,
400
- reduction=reduction,
401
- softcap=softcap,
402
- RETURN_Z_LOSS=return_z_loss,
403
- RETURN_TOKEN_ACCURACY=return_token_accuracy,
404
- RETURN_PREDICTED_TOKENS=return_predicted_tokens,
405
- BLOCK_SIZE=BLOCK_SIZE,
406
- HAS_WEIGHT=True if weight is not None else False,
407
- HAS_SOFTCAPPING=True if softcap is not None else False,
408
- HAS_GRADIENTS=_input.requires_grad,
409
- # TODO: 32 seems to give the best performance
410
- # Performance is quite sensitive to num_warps
411
- num_warps=32 if not is_hip() else 16,
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
- BT, V = _input.shape
441
- n_rows = BT
442
- BLOCK_SIZE = min(MAX_FUSED_SIZE, triton.next_power_of_2(V))
443
-
444
- element_mul_kernel[(n_rows,)](
445
- _input,
446
- _input.stride(-2),
447
- grad_output,
448
- V,
449
- BLOCK_SIZE=BLOCK_SIZE,
450
- num_warps=32 if not is_hip() else 16,
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
- grid = lambda meta: (triton.cdiv(N, meta["BLOCK_N"]), M)
111
- _dyt_fwd_kernel[grid](
112
- x,
113
- y,
114
- alpha,
115
- gamma,
116
- beta,
117
- HAVE_BETA,
118
- N,
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
- grid = lambda meta: (triton.cdiv(N, meta["BLOCK_N"]), NUM_SMS)
143
- _dyt_bwd_kernel[grid](dy, dx, da, dg, db, x, alpha, gamma, HAVE_BETA, M, N)
 
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
- # Here we calculate the gradient of logits_chunk in place so we can save memory.
151
- liger_cross_entropy_kernel[(n_rows,)](
152
- X_ptr=logits_chunk,
153
- X_stride=logits_chunk.stride(-2),
154
- Y_ptr=target_chunk,
155
- Y_stride=target_chunk.stride(-1), # always 1
156
- weight_ptr=ce_weight,
157
- loss_ptr=loss_1d_slice,
158
- z_loss_ptr=z_loss_1d_slice,
159
- loss_stride=loss_1d_slice.stride(-1), # always 1
160
- token_accuracy_ptr=token_accuracy_1d_slice,
161
- token_accuracy_stride=token_accuracy_1d_slice.stride(-1)
162
- if return_token_accuracy
163
- else 0, # always 1 if accuracy is enabled
164
- predicted_tokens_ptr=predicted_tokens_1d_slice,
165
- predicted_tokens_stride=predicted_tokens_1d_slice.stride(-1)
166
- if return_predicted_tokens
167
- else 0, # always 1 if predicted tokens is enabled
168
- n_cols=V,
169
- n_non_ignore=total_n_non_ignore,
170
- sum_non_ignore_weight=total_sum_non_ignore_ce_weight,
171
- weight_sum=ce_weight_sum,
172
- ignore_index=ignore_index,
173
- lse_square_scale=lse_square_scale,
174
- label_smoothing=label_smoothing,
175
- reduction=reduction,
176
- softcap=softcap,
177
- RETURN_Z_LOSS=return_z_loss,
178
- RETURN_TOKEN_ACCURACY=return_token_accuracy,
179
- RETURN_PREDICTED_TOKENS=return_predicted_tokens,
180
- HAS_WEIGHT=True if ce_weight is not None else False,
181
- HAS_SOFTCAPPING=True if softcap is not None else False,
182
- HAS_GRADIENTS=input_requires_grad,
183
- BLOCK_SIZE=BLOCK_SIZE,
184
- num_warps=32 if not is_hip() else 16,
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
- # We use a Triton kernel instead of a PyTorch operation because modifying inputs in-place
251
- # for gradient storage and backward multiple times causes anomalies with PyTorch but not with Triton.
252
- BT, H = grad_input.shape
253
- n_rows = BT
254
- BLOCK_SIZE = min(MAX_FUSED_SIZE, triton.next_power_of_2(H))
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
- grad_weight,
272
- grad_weight.stride(-2),
273
  grad_output,
274
  H,
275
  BLOCK_SIZE=BLOCK_SIZE,
276
  num_warps=32 if not is_hip() else 16,
277
  )
278
 
279
- if grad_bias is not None:
280
- V = grad_bias.shape[0]
281
- n_rows = V
282
-
283
- element_mul_kernel[(n_rows,)](
284
- grad_bias,
285
- grad_bias.stride(-1),
286
- grad_output,
287
- 1,
288
- BLOCK_SIZE=BLOCK_SIZE,
289
- num_warps=32 if not is_hip() else 16,
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
- _geglu_tanh_forward_kernel[(n_rows,)](
98
- a,
99
- b,
100
- c,
101
- c.stride(-2),
102
- n_cols=n_cols,
103
- BLOCK_SIZE=BLOCK_SIZE,
104
- num_warps=num_warps,
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
- _geglu_tanh_backward_kernel[(n_rows,)](
118
- dc,
119
- a,
120
- b,
121
- dc.stride(-2),
122
- n_cols=n_cols,
123
- BLOCK_SIZE=BLOCK_SIZE,
124
- num_warps=num_warps,
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
- _group_norm_forward_kernel[(batch_size, num_groups)](
219
- Y,
220
- Y.stride(0),
221
- Y.stride(1),
222
- X,
223
- X.stride(0),
224
- X.stride(1),
225
- Mean,
226
- Mean.stride(0),
227
- Mean.stride(1),
228
- RSTD,
229
- RSTD.stride(0),
230
- RSTD.stride(1),
231
- W,
232
- B,
233
- hidden_size,
234
- channels_per_group,
235
- eps,
236
- BLOCK_SIZE=BLOCK_SIZE,
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
- BLOCK_SIZE = min(MAX_FUSED_SIZE, triton.next_power_of_2(hidden_size))
258
- _group_norm_backward_kernel[(batch_size, num_groups)](
259
- X,
260
- X.stride(0),
261
- X.stride(1),
262
- W,
263
- Mean,
264
- Mean.stride(0),
265
- Mean.stride(1),
266
- RSTD,
267
- DX,
268
- DW,
269
- DB,
270
- dY,
271
- hidden_size,
272
- channels_per_group,
273
- BLOCK_SIZE=BLOCK_SIZE,
274
- dtype=triton_dtype,
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
- _jsd_kernel[(n_rows,)](
113
- X_ptr=_input, # input in logspace, X = log Q
114
- X_stride=_input.stride(-2),
115
- Y_ptr=target, # ground truth in logspace, Y = log P
116
- Y_stride=target.stride(-2),
117
- loss_ptr=loss,
118
- loss_stride=loss.stride(-2),
119
- dX_ptr=dX,
120
- dX_stride=dX.stride(-2),
121
- label_ptr=(shift_labels if has_label else torch.empty(1, device=_input.device)), # dummy ptr if no label
122
- beta=beta,
123
- n_non_ignore=n_non_ignore,
124
- ignore_index=ignore_index,
125
- n_cols=V,
126
- BLOCK_SIZE=BLOCK_SIZE,
127
- HAS_LABEL=has_label,
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
- _kldiv_kernel_forward[grid](
134
- y_pred,
135
- y_pred.stride(0),
136
- y_true,
137
- y_true.stride(0),
138
- output_tensor,
139
- output_tensor.stride(0),
140
- V,
141
- eps=eps,
142
- BLOCK_SIZE=BLOCK_SIZE,
143
- num_warps=num_warps,
144
- log_target=log_target,
145
- reduction=reduction,
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
- # We store the gradients in-place in the input tensor
169
- _kldiv_kernel_backward[grid](
170
- target,
171
- target.stride(0),
172
- new_grads,
173
- new_grads.stride(0),
174
- V,
175
- BLOCK_SIZE=BLOCK_SIZE,
176
- num_warps=num_warps,
177
- log_target=log_target,
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
- # Launch kernel with one thread block per row for optimal performance
206
- grid = (n_rows,)
207
- _layer_norm_forward_kernel[grid](
208
- Y,
209
- Y.stride(0),
210
- X,
211
- X.stride(0),
212
- W,
213
- W.stride(0),
214
- B,
215
- B.stride(0),
216
- Mean,
217
- Mean.stride(0),
218
- RSTD,
219
- RSTD.stride(0),
220
- n_cols,
221
- eps,
222
- BLOCK_SIZE=BLOCK_SIZE,
223
- num_warps=num_warps,
224
- **kernel_args,
 
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
- # Launch kernel with one thread block per row for optimal performance
277
- _layer_norm_backward_kernel[grid](
278
- X,
279
- X.stride(0),
280
- W,
281
- Mean,
282
- Mean.stride(0),
283
- RSTD,
284
- RSTD.stride(0),
285
- DX,
286
- DX.stride(0),
287
- _DW,
288
- _DW.stride(0),
289
- _DB,
290
- _DB.stride(0),
291
- dY,
292
- dY.stride(0),
293
- n_rows,
294
- n_cols,
295
- rows_per_program=rows_per_program,
296
- BLOCK_SIZE=BLOCK_SIZE,
297
- **kernel_args,
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": "_liger_kernels_rocm_4d9f798",
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": "pdeXrc8M9KCvCpe0mtloEEFT3iJWZab9CMNsgT2CK10=",
15
- "cross_entropy.py": "LsQN0OI8VHqgUObWF0qqTgNA9z7SYMaZNb2I7FIYnO4=",
16
- "dyt.py": "806Tai0T+Yi7Ctf5YGmtwAc87MUjlzjsLky7b1T+oWM=",
17
- "fused_linear_cross_entropy.py": "yjtZHWyBcXCBJ57NbbrQ1wMmTr9nARqDBrdjXVYpR74=",
18
- "geglu.py": "Lmw0I38IBIHkNROTf9V3bP6PZLktCWCqNAxzVRaipbk=",
19
- "group_norm.py": "xdY/tCr4rHOb4PhD827259XB7QDb22z2psUKzTYTKBE=",
20
- "jsd.py": "CftyIN36AU63e1GdtuA7hNhs66I3TjS9lttcFriBnGA=",
21
- "kl_div.py": "VQG/89mefLifPeVEv+1O8hmsXgGujMLA0JXDwO5QvtU=",
22
- "layer_norm.py": "INpxK7Ac27f+eckcMfOOSgZErzDhNF5iGt207eV9aoM=",
23
  "layers.py": "+50F9xmnpQXbU1K/lZxdJwbnTotkksJlhoLOWHjo+/M=",
24
  "liger_kernels/__init__.py": "DFYPlrhXwYjEqCl/8n0SmWGZV8NFml5DPhMjKfv98GY=",
25
- "qwen2vl_mrope.py": "3GExhYpLgB4VUtyZyjRk8XjEur3W4EWF6HQ67ML5vBU=",
26
- "rms_norm.py": "GLzIZjDP74QxaeKQ5B/VDXKTGxD6Xo9cVaW/iqmsc4c=",
27
- "rope.py": "v+7JHRrv+5ImoROkpKfl30WwWI4qTa2tAl7zQeB4ml4=",
28
- "swiglu.py": "TLhyvpr7nbVbMqWbx8uz+iU2GxS6q0QrVxWAd1eQCBo=",
29
  "tiled_mlp.py": "PJw2R8YdHHxQRSwbFJWjcHxgzssEqsRWqyogkY8pKR8=",
30
- "tvd.py": "VbWUtpJ79E2atqjBoVyTSO785aS+idHiQBIrJwifDAw=",
31
- "utils.py": "Y8wQsplz6ed7oPmBfX7yGGcX74VB3t/e5VFvYv/X7wU="
32
  }
33
  },
34
  "provenance": {
@@ -38,7 +38,7 @@
38
  "dirty": false
39
  },
40
  "kernel": {
41
- "sha": "4d9f798b95ed280b770213048b8d3f41509d0c9e",
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
- _triton_qwen2vl_mrope[(n_row,)](
132
- q,
133
- k,
134
- cos,
135
- sin,
136
- seq_len,
137
- batch_size,
138
- n_q_head,
139
- n_kv_head,
140
- head_dim,
141
- pad_n_q_head,
142
- pad_n_kv_head,
143
- pad_hd,
144
- mrope_section[0],
145
- mrope_section[1],
146
- BLOCK_SIZE=BLOCK_SIZE,
147
- BACKWARD_PASS=False,
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
- # backward is similar to forward except swapping few ops
170
- _triton_qwen2vl_mrope[(n_row,)](
171
- dq,
172
- dk,
173
- cos,
174
- sin,
175
- seq_len,
176
- batch_size,
177
- n_q_head,
178
- n_kv_head,
179
- head_dim,
180
- pad_n_q_head,
181
- pad_n_kv_head,
182
- pad_hd,
183
- mrope_section[0],
184
- mrope_section[1],
185
- BLOCK_SIZE=BLOCK_SIZE,
186
- BACKWARD_PASS=True,
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
- if BLOCK_SIZE > 256 or n_rows < 4096 * 8 or row_mode:
442
- _rms_norm_forward_kernel[(n_rows,)](
443
- Y,
444
- Y.stride(0),
445
- X,
446
- X.stride(0),
447
- W,
448
- W.stride(0) if elementwise_affine else 0,
449
- RSTD,
450
- RSTD.stride(0),
451
- n_cols,
452
- eps,
453
- offset,
454
- casting_mode,
455
- elementwise_affine=elementwise_affine,
456
- BLOCK_SIZE=BLOCK_SIZE,
457
- num_warps=num_warps,
458
- **kernel_args, # XPU-specific optimization
459
- )
460
- else:
461
- BLOCK_ROW = 16
462
- kernel_args["BLOCK_ROW"] = BLOCK_ROW
463
- _block_rms_norm_forward_kernel[(triton.cdiv(n_rows, BLOCK_ROW),)](
464
- Y,
465
- Y.stride(0),
466
- X,
467
- X.stride(0),
468
- W,
469
- W.stride(0) if elementwise_affine else 0,
470
- RSTD,
471
- RSTD.stride(0),
472
- n_rows,
473
- n_cols,
474
- eps,
475
- offset,
476
- casting_mode,
477
- elementwise_affine=elementwise_affine,
478
- BLOCK_SIZE=BLOCK_SIZE,
479
- num_warps=num_warps,
480
- **kernel_args, # XPU-specific optimization
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
- if BLOCK_SIZE > 256 or n_rows < 4096 * 8 or row_mode:
523
- _rms_norm_backward_kernel[grid](
524
- dY,
525
- dY.stride(0),
526
- dX,
527
- dX.stride(0),
528
- X,
529
- X.stride(0),
530
- torch_to_triton_dtype[X.dtype],
531
- W,
532
- W.stride(0) if elementwise_affine else 0,
533
- RSTD,
534
- RSTD.stride(0),
535
- _dW,
536
- _dW.stride(0) if elementwise_affine else 0,
537
- n_rows,
538
- n_cols,
539
- offset,
540
- rows_per_program,
541
- casting_mode,
542
- elementwise_affine=elementwise_affine,
543
- BLOCK_SIZE=BLOCK_SIZE,
544
- num_warps=num_warps,
545
- **kernel_args, # XPU-specific optimization
546
- )
547
- else:
548
- BLOCK_ROW = 16
549
- kernel_args["BLOCK_ROW"] = BLOCK_ROW
550
- _block_rms_norm_backward_kernel[grid](
551
- dY,
552
- dY.stride(0),
553
- dX,
554
- dX.stride(0),
555
- X,
556
- X.stride(0),
557
- torch_to_triton_dtype[X.dtype],
558
- W,
559
- W.stride(0) if elementwise_affine else 0,
560
- RSTD,
561
- RSTD.stride(0),
562
- _dW,
563
- _dW.stride(0) if elementwise_affine else 0,
564
- n_rows,
565
- n_cols,
566
- offset,
567
- casting_mode,
568
- elementwise_affine=elementwise_affine,
569
- BLOCK_SIZE=BLOCK_SIZE,
570
- num_warps=num_warps,
571
- **kernel_args, # XPU-specific optimization
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
- _triton_rope[(n_row,)](
138
- q,
139
- q.stride(1),
140
- k,
141
- k.stride(1),
142
- cos,
143
- cos.stride(-2),
144
- sin,
145
- sin.stride(-2),
146
- seq_len,
147
- batch_size,
148
- cos_batch_size,
149
- n_q_head,
150
- n_kv_head,
151
- head_dim,
152
- pad_n_q_head,
153
- pad_n_kv_head,
154
- pad_hd,
155
- BLOCK_SIZE=BLOCK_SIZE,
156
- BACKWARD_PASS=False,
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
- # backward is similar to forward except swapping few ops
180
- _triton_rope[(n_row,)](
181
- dq,
182
- dq.stride(1),
183
- dk,
184
- dk.stride(1),
185
- cos,
186
- cos.stride(-2),
187
- sin,
188
- sin.stride(-2),
189
- seq_len,
190
- batch_size,
191
- cos_batch_size,
192
- n_q_head,
193
- n_kv_head,
194
- head_dim,
195
- pad_n_q_head,
196
- pad_n_kv_head,
197
- pad_hd,
198
- BLOCK_SIZE=BLOCK_SIZE,
199
- BACKWARD_PASS=True,
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
- _swiglu_forward_kernel[(n_rows,)](
77
- a,
78
- b,
79
- c,
80
- c.stride(-2),
81
- float(gate_multiplier),
82
- n_cols=n_cols,
83
- BLOCK_SIZE=BLOCK_SIZE,
84
- num_warps=num_warps,
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
- _swiglu_backward_kernel[(n_rows,)](
98
- dc,
99
- a,
100
- b,
101
- dc.stride(-2),
102
- float(gate_multiplier),
103
- n_cols=n_cols,
104
- BLOCK_SIZE=BLOCK_SIZE,
105
- num_warps=num_warps,
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
- _tv_distance_kernel[grid](
128
- p,
129
- p.stride(0),
130
- q,
131
- q.stride(0),
132
- output_tensor,
133
- output_tensor.stride(0),
134
- grads,
135
- grads.stride(0),
136
- shift_labels if has_label else torch.empty(1, device=p.device),
137
- ignore_index,
138
- V,
139
- scale,
140
- BLOCK_SIZE=BLOCK_SIZE,
141
- HAS_LABEL=has_label,
142
- num_warps=num_warps,
143
- reduction=reduction,
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