Pytorch:C 语言扩展 pytorch

pytorch利用 CFFI 进行 C 语言扩展。包括两个基本的步骤(docs):

  1. 编写 C 代码;
  2. python 调用 C 代码,实现相应的 Function 或 Module。

C语言拓展需要上文提到的torch.util.ffi模块,比如官网给出的定义一个加法运算(一个栗子),这里以定义ReLU函数为例:

1. C 代码

pytorch C 的基本数据结构是 THTensor(THFloatTensor、THByteTensor等)。我们以简单的 ReLU 函数为例,示例编写 C 。

y=ReLU(x)=max(x,0)

Function 需要定义前向和后向两个方向的操作,因此,C 代码要实现相应的功能。

1.1 头文件 ext_lib.h

/* ext_lib.h */
int relu_forward(THFloatTensor *input, THFloatTensor *output);
int relu_backward(THFloatTensor *grad_output, THFloatTensor *input, THFloatTensor *grad_input);

1.1 函数实现 ext_lib.c

/* ext_lib.c */

#include <TH/TH.h>

int relu_forward(THFloatTensor *input, THFloatTensor *output)
{
  THFloatTensor_resizeAs(output, input);
  THFloatTensor_clamp(output, input, 0, INFINITY);
  return 1;
}

int relu_backward(THFloatTensor *grad_output, THFloatTensor *input, THFloatTensor *grad_input)
{
  THFloatTensor_resizeAs(grad_input, grad_output);
  THFloatTensor_zero(grad_input);

  THLongStorage* size = THFloatTensor_newSizeOf(grad_output);
  THLongStorage *stride = THFloatTensor_newStrideOf(grad_output);
  THByteTensor *mask = THByteTensor_newWithSize(size, stride);

  THFloatTensor_geValue(mask, input, 0);
  THFloatTensor_maskedCopy(grad_input, mask, grad_output);
  return 1;
}

 

评论
添加红包

请填写红包祝福语或标题

红包个数最小为10个

红包金额最低5元

当前余额3.43前往充值 >
需支付:10.00
成就一亿技术人!
领取后你会自动成为博主和红包主的粉丝 规则
hope_wisdom
发出的红包
实付
使用余额支付
点击重新获取
扫码支付
钱包余额 0

抵扣说明:

1.余额是钱包充值的虚拟货币,按照1:1的比例进行支付金额的抵扣。
2.余额无法直接购买下载,可以购买VIP、付费专栏及课程。

余额充值