Skip to content
体验新版
项目
组织
正在加载...
登录
切换导航
打开侧边栏
BaiXuePrincess
Paddle
提交
981fc2bd
P
Paddle
项目概览
BaiXuePrincess
/
Paddle
与 Fork 源项目一致
Fork自
PaddlePaddle / Paddle
通知
1
Star
1
Fork
0
代码
文件
提交
分支
Tags
贡献者
分支图
Diff
Issue
0
列表
看板
标记
里程碑
合并请求
0
Wiki
0
Wiki
分析
仓库
DevOps
项目成员
Pages
P
Paddle
项目概览
项目概览
详情
发布
仓库
仓库
文件
提交
分支
标签
贡献者
分支图
比较
Issue
0
Issue
0
列表
看板
标记
里程碑
合并请求
0
合并请求
0
Pages
分析
分析
仓库分析
DevOps
Wiki
0
Wiki
成员
成员
收起侧边栏
关闭侧边栏
动态
分支图
创建新Issue
提交
Issue看板
未验证
提交
981fc2bd
编写于
1月 25, 2019
作者:
T
tangwei12
提交者:
GitHub
1月 25, 2019
浏览文件
操作
浏览文件
下载
电子邮件补丁
差异文件
fix bug in merge_ids (#15503)
* fix mistakes in merge_ids, test=develop
上级
a7ba07d7
变更
1
隐藏空白更改
内联
并排
Showing
1 changed file
with
9 addition
and
9 deletion
+9
-9
paddle/fluid/operators/distributed_ops/merge_ids_op.h
paddle/fluid/operators/distributed_ops/merge_ids_op.h
+9
-9
未找到文件。
paddle/fluid/operators/distributed_ops/merge_ids_op.h
浏览文件 @
981fc2bd
...
@@ -43,9 +43,9 @@ class MergeIdsOpKernel : public framework::OpKernel<T> {
...
@@ -43,9 +43,9 @@ class MergeIdsOpKernel : public framework::OpKernel<T> {
PADDLE_ENFORCE_EQ
(
ids
.
size
(),
outs
.
size
(),
PADDLE_ENFORCE_EQ
(
ids
.
size
(),
outs
.
size
(),
"the number of Ids and Out should be the same"
);
"the number of Ids and Out should be the same"
);
size
_t
row_ids_size
=
0
;
int64
_t
row_ids_size
=
0
;
int
row_size
=
0
;
int
64_t
row_size
=
0
;
int
embedding_size
=
0
;
int
64_t
embedding_size
=
0
;
for
(
size_t
i
=
0
;
i
<
x_tensors
.
size
();
++
i
)
{
for
(
size_t
i
=
0
;
i
<
x_tensors
.
size
();
++
i
)
{
const
auto
*
x_tensor
=
x_tensors
[
i
];
const
auto
*
x_tensor
=
x_tensors
[
i
];
...
@@ -69,7 +69,7 @@ class MergeIdsOpKernel : public framework::OpKernel<T> {
...
@@ -69,7 +69,7 @@ class MergeIdsOpKernel : public framework::OpKernel<T> {
for
(
size_t
i
=
0
;
i
<
x_tensors
.
size
();
++
i
)
{
for
(
size_t
i
=
0
;
i
<
x_tensors
.
size
();
++
i
)
{
const
auto
*
row_id
=
row_ids
[
i
];
const
auto
*
row_id
=
row_ids
[
i
];
for
(
int
j
=
0
;
j
<
row_id
->
numel
();
++
j
)
{
for
(
auto
j
=
0
;
j
<
row_id
->
numel
();
++
j
)
{
int64_t
key
=
row_id
->
data
<
int64_t
>
()[
j
];
int64_t
key
=
row_id
->
data
<
int64_t
>
()[
j
];
std
::
tuple
<
int64_t
,
int64_t
>
val
=
std
::
make_tuple
(
i
,
j
);
std
::
tuple
<
int64_t
,
int64_t
>
val
=
std
::
make_tuple
(
i
,
j
);
selected_rows_idx_map
.
insert
(
std
::
make_pair
(
key
,
val
));
selected_rows_idx_map
.
insert
(
std
::
make_pair
(
key
,
val
));
...
@@ -84,13 +84,13 @@ class MergeIdsOpKernel : public framework::OpKernel<T> {
...
@@ -84,13 +84,13 @@ class MergeIdsOpKernel : public framework::OpKernel<T> {
out
->
set_lod
(
out_ids
->
lod
());
out
->
set_lod
(
out_ids
->
lod
());
int
nums
=
static_cast
<
int
>
(
out_ids
->
dims
()[
0
])
;
auto
nums
=
out_ids
->
dims
()[
0
]
;
auto
*
out_data
=
out
->
mutable_data
<
T
>
(
auto
*
out_data
=
out
->
mutable_data
<
T
>
(
framework
::
make_ddim
({
nums
,
embedding_size
}),
place
);
framework
::
make_ddim
({
nums
,
embedding_size
}),
place
);
for
(
int
j
=
0
;
j
<
nums
;
++
j
)
{
for
(
auto
j
=
0
;
j
<
nums
;
++
j
)
{
int
id
=
out_ids
->
data
<
int64_t
>
()[
j
];
auto
id
=
out_ids
->
data
<
int64_t
>
()[
j
];
auto
row_tuple
=
selected_rows_idx_map
[
id
]
;
auto
row_tuple
=
selected_rows_idx_map
.
at
(
id
)
;
int64_t
row_idx
=
std
::
get
<
1
>
(
row_tuple
);
auto
row_idx
=
std
::
get
<
1
>
(
row_tuple
);
const
auto
*
x_tensor
=
x_tensors
[
std
::
get
<
0
>
(
row_tuple
)];
const
auto
*
x_tensor
=
x_tensors
[
std
::
get
<
0
>
(
row_tuple
)];
memcpy
(
out_data
+
embedding_size
*
j
,
memcpy
(
out_data
+
embedding_size
*
j
,
...
...
编辑
预览
Markdown
is supported
0%
请重试
或
添加新附件
.
添加附件
取消
You are about to add
0
people
to the discussion. Proceed with caution.
先完成此消息的编辑!
取消
想要评论请
注册
或
登录