@@ -20,13 +20,13 @@ void remove_padding(int64_t *output_data,
2020 const int *cum_offsets,
2121 const int sequence_length,
2222 const int bsz) {
23- for (int bi = 0 ; bi < bsz; ++bi) {
24- for (int i = 0 ; i < seq_lens[bi]; ++i) {
25- const int tgt_seq_id = bi * sequence_length - cum_offsets[bi] + i;
26- const int src_seq_id = bi * sequence_length + i;
27- output_data[tgt_seq_id] = input_data[src_seq_id];
28- }
23+ for (int bi = 0 ; bi < bsz; ++bi) {
24+ for (int i = 0 ; i < seq_lens[bi]; ++i) {
25+ const int tgt_seq_id = bi * sequence_length - cum_offsets[bi] + i;
26+ const int src_seq_id = bi * sequence_length + i;
27+ output_data[tgt_seq_id] = input_data[src_seq_id];
2928 }
29+ }
3030}
3131
3232void get_padding_offset_kernel (int *padding_offset,
@@ -37,85 +37,77 @@ void get_padding_offset_kernel(int *padding_offset,
3737 const int *seq_lens,
3838 const int max_seq_len,
3939 const int bsz) {
40- for (int bi = 0 ; bi < bsz; ++bi) {
41- int cum_offset = bi == 0 ? 0 : cum_offsets[bi - 1 ];
42- auto seq_len_now = seq_lens[bi];
43- for (int i = 0 ; i < seq_len_now; ++i) {
44- padding_offset[bi * max_seq_len - cum_offset + i] = cum_offset;
45- }
46- cum_offsets_out[bi] = cum_offset;
47- int cum_seq_len = (bi + 1 ) * max_seq_len - cum_offsets[bi];
48- cu_seqlens_q[bi + 1 ] = cum_seq_len;
49- cu_seqlens_k[bi + 1 ] = cum_seq_len;
40+ for (int bi = 0 ; bi < bsz; ++bi) {
41+ int cum_offset = bi == 0 ? 0 : cum_offsets[bi - 1 ];
42+ auto seq_len_now = seq_lens[bi];
43+ for (int i = 0 ; i < seq_len_now; ++i) {
44+ padding_offset[bi * max_seq_len - cum_offset + i] = cum_offset;
5045 }
46+ cum_offsets_out[bi] = cum_offset;
47+ int cum_seq_len = (bi + 1 ) * max_seq_len - cum_offsets[bi];
48+ cu_seqlens_q[bi + 1 ] = cum_seq_len;
49+ cu_seqlens_k[bi + 1 ] = cum_seq_len;
50+ }
5151}
5252
5353std::vector<paddle::Tensor> GetPaddingOffset (const paddle::Tensor &input_ids,
5454 const paddle::Tensor &cum_offsets,
5555 const paddle::Tensor &token_num,
5656 const paddle::Tensor &seq_len) {
57- std::vector<int64_t > input_ids_shape = input_ids.shape ();
58- const int bsz = seq_len.shape ()[0 ];
59- const int seq_length = input_ids_shape[1 ];
60- auto cum_offsets_out = cum_offsets.copy_to (paddle::CPUPlace (), false );
61- auto cpu_token_num = token_num.copy_to (paddle::CPUPlace (), false );
57+ std::vector<int64_t > input_ids_shape = input_ids.shape ();
58+ const int bsz = seq_len.shape ()[0 ];
59+ const int seq_length = input_ids_shape[1 ];
60+ auto cum_offsets_out = cum_offsets.copy_to (paddle::CPUPlace (), false );
61+ auto cpu_token_num = token_num.copy_to (paddle::CPUPlace (), false );
6262
63- const int token_num_data = cpu_token_num.data <int64_t >()[0 ];
64- auto x_remove_padding = paddle::empty (
65- {token_num_data}, paddle::DataType::INT64 , input_ids.place ());
66- auto padding_offset = paddle::empty (
67- {token_num_data}, paddle::DataType::INT32 , input_ids.place ());
68- auto cu_seqlens_q =
69- paddle::full ({bsz + 1 }, 0 , paddle::DataType::INT32 , input_ids.place ());
70- auto cu_seqlens_k =
71- paddle::full ({bsz + 1 }, 0 , paddle::DataType::INT32 , input_ids.place ());
72- get_padding_offset_kernel (padding_offset.data <int >(),
73- cum_offsets_out.data <int >(),
74- cu_seqlens_q.data <int >(),
75- cu_seqlens_k.data <int >(),
76- cum_offsets.data <int >(),
77- seq_len.data <int >(),
78- seq_length,
79- bsz);
80- remove_padding (x_remove_padding.data <int64_t >(),
81- input_ids.data <int64_t >(),
82- seq_len.data <int >(),
83- cum_offsets_out.data <int >(),
84- seq_length,
85- bsz);
86- return {x_remove_padding,
87- padding_offset,
88- cu_seqlens_q,
89- cu_seqlens_k};
63+ const int token_num_data = cpu_token_num.data <int64_t >()[0 ];
64+ auto x_remove_padding = paddle::empty (
65+ {token_num_data}, paddle::DataType::INT64 , input_ids.place ());
66+ auto padding_offset = paddle::empty (
67+ {token_num_data}, paddle::DataType::INT32 , input_ids.place ());
68+ auto cu_seqlens_q =
69+ paddle::full ({bsz + 1 }, 0 , paddle::DataType::INT32 , input_ids.place ());
70+ auto cu_seqlens_k =
71+ paddle::full ({bsz + 1 }, 0 , paddle::DataType::INT32 , input_ids.place ());
72+ get_padding_offset_kernel (padding_offset.data <int >(),
73+ cum_offsets_out.data <int >(),
74+ cu_seqlens_q.data <int >(),
75+ cu_seqlens_k.data <int >(),
76+ cum_offsets.data <int >(),
77+ seq_len.data <int >(),
78+ seq_length,
79+ bsz);
80+ remove_padding (x_remove_padding.data <int64_t >(),
81+ input_ids.data <int64_t >(),
82+ seq_len.data <int >(),
83+ cum_offsets_out.data <int >(),
84+ seq_length,
85+ bsz);
86+ return {x_remove_padding, padding_offset, cu_seqlens_q, cu_seqlens_k};
9087}
9188
9289std::vector<std::vector<int64_t >> GetPaddingOffsetInferShape (
9390 const std::vector<int64_t > &input_ids_shape,
9491 const std::vector<int64_t > &cum_offsets_shape,
9592 const std::vector<int64_t > &token_num_shape,
9693 const std::vector<int64_t > &seq_len_shape) {
97- int64_t bsz = seq_len_shape[0 ];
98- int64_t seq_len = input_ids_shape[1 ];
99- return {{-1 }, {-1 }, {bsz + 1 }, {bsz + 1 }};
94+ int64_t bsz = seq_len_shape[0 ];
95+ int64_t seq_len = input_ids_shape[1 ];
96+ return {{-1 }, {-1 }, {bsz + 1 }, {bsz + 1 }};
10097}
10198
10299std::vector<paddle::DataType> GetPaddingOffsetInferDtype (
103100 const paddle::DataType &input_ids_dtype,
104101 const paddle::DataType &cum_offsets_dtype,
105102 const paddle::DataType &token_num_dtype,
106103 const paddle::DataType &seq_len_dtype) {
107- return {input_ids_dtype,
108- seq_len_dtype,
109- seq_len_dtype,
110- seq_len_dtype};
104+ return {input_ids_dtype, seq_len_dtype, seq_len_dtype, seq_len_dtype};
111105}
112106
113107PD_BUILD_STATIC_OP (get_padding_offset_cpu)
114108 .Inputs({" input_ids" , " cum_offsets" , " token_num" , " seq_len" })
115- .Outputs({" x_remove_padding" ,
116- " padding_offset" ,
117- " cu_seqlens_q" ,
118- " cu_seqlens_k" })
109+ .Outputs(
110+ {" x_remove_padding" , " padding_offset" , " cu_seqlens_q" , " cu_seqlens_k" })
119111 .SetKernelFn(PD_KERNEL (GetPaddingOffset))
120112 .SetInferShapeFn(PD_INFER_SHAPE (GetPaddingOffsetInferShape))
121113 .SetInferDtypeFn(PD_INFER_DTYPE (GetPaddingOffsetInferDtype));
0 commit comments