@@ -528,11 +528,16 @@ void ggml_cuda_op_rope_impl(ggml_backend_cuda_context & ctx,
528528 }
529529 cudaStream_t stream = ctx.stream ();
530530
531- GGML_ASSERT (src0->type == GGML_TYPE_F32 || src0->type == GGML_TYPE_F16 );
532- GGML_ASSERT ( dst->type == GGML_TYPE_F32 || dst->type == GGML_TYPE_F16 );
531+ // [P1 bf16-stream] bf16 is accepted so a bf16-resident attention segment can rope
532+ // its Q/K in bf16 (the norm/neox kernels are float-internal, so the bf16 arms just
533+ // add T/D = nv_bfloat16 instantiations below).
534+ GGML_ASSERT (src0->type == GGML_TYPE_F32 || src0->type == GGML_TYPE_F16 || src0->type == GGML_TYPE_BF16 );
535+ GGML_ASSERT ( dst->type == GGML_TYPE_F32 || dst->type == GGML_TYPE_F16 || dst->type == GGML_TYPE_BF16 );
533536 // When not fused, src0 and dst types must match
534537 // When fused (ROPE+VIEW+SET_ROWS), src0 may be F32 and dst may be F16
535- GGML_ASSERT (src0->type == dst->type || (src0->type == GGML_TYPE_F32 && dst->type == GGML_TYPE_F16 ));
538+ GGML_ASSERT (src0->type == dst->type ||
539+ (src0->type == GGML_TYPE_F32 && dst->type == GGML_TYPE_F16 ) ||
540+ (src0->type == GGML_TYPE_F32 && dst->type == GGML_TYPE_BF16 ));
536541
537542 const int64_t ne00 = src0->ne [0 ]; // head dims
538543 const int64_t ne01 = src0->ne [1 ]; // num heads
@@ -610,6 +615,16 @@ void ggml_cuda_op_rope_impl(ggml_backend_cuda_context & ctx,
610615 s03, s1, s2, s3, n_dims, nr, pos, freq_scale, freq_base,
611616 ext_factor, attn_factor, corr_dims, freq_factors, row_indices,
612617 set_rows_stride, stream);
618+ } else if (src0->type == GGML_TYPE_BF16 && dst_type == GGML_TYPE_BF16 ) {
619+ rope_neox_cuda<forward, nv_bfloat16, nv_bfloat16>((const nv_bfloat16 *) src0_d, (nv_bfloat16 *) dst_d, ne00, ne01, ne02, s01, s02,
620+ s03, s1, s2, s3, n_dims, nr, pos, freq_scale, freq_base,
621+ ext_factor, attn_factor, corr_dims, freq_factors, row_indices,
622+ set_rows_stride, stream);
623+ } else if (src0->type == GGML_TYPE_F32 && dst_type == GGML_TYPE_BF16 ) {
624+ rope_neox_cuda<forward, float , nv_bfloat16>((const float *) src0_d, (nv_bfloat16 *) dst_d, ne00, ne01, ne02, s01, s02,
625+ s03, s1, s2, s3, n_dims, nr, pos, freq_scale, freq_base,
626+ ext_factor, attn_factor, corr_dims, freq_factors, row_indices,
627+ set_rows_stride, stream);
613628 } else {
614629 GGML_ABORT (" fatal error" );
615630 }
@@ -653,6 +668,16 @@ void ggml_cuda_op_rope_impl(ggml_backend_cuda_context & ctx,
653668 s03, s1, s2, s3, n_dims, nr, pos, freq_scale, freq_base,
654669 ext_factor, attn_factor, corr_dims, freq_factors, row_indices,
655670 set_rows_stride, stream);
671+ } else if (src0->type == GGML_TYPE_BF16 && dst_type == GGML_TYPE_BF16 ) {
672+ rope_norm_cuda<forward, nv_bfloat16, nv_bfloat16>((const nv_bfloat16 *) src0_d, (nv_bfloat16 *) dst_d, ne00, ne01, ne02, s01, s02,
673+ s03, s1, s2, s3, n_dims, nr, pos, freq_scale, freq_base,
674+ ext_factor, attn_factor, corr_dims, freq_factors, row_indices,
675+ set_rows_stride, stream);
676+ } else if (src0->type == GGML_TYPE_F32 && dst_type == GGML_TYPE_BF16 ) {
677+ rope_norm_cuda<forward, float , nv_bfloat16>((const float *) src0_d, (nv_bfloat16 *) dst_d, ne00, ne01, ne02, s01, s02,
678+ s03, s1, s2, s3, n_dims, nr, pos, freq_scale, freq_base,
679+ ext_factor, attn_factor, corr_dims, freq_factors, row_indices,
680+ set_rows_stride, stream);
656681 } else {
657682 GGML_ABORT (" fatal error" );
658683 }
0 commit comments