CNN卷积神经网络Python实现

python 复制代码
import torch
from torch import nn

# ①定义互相关运算
def corr2d(X, K):
    """计算二维互相关运算。"""
    # 获取K的形状 行为h,列为w
    h, w = K.shape
    # 生成全0的矩阵,行为X的行减去h加上1,列为X的列减去w加上1
    Y = torch.zeros((X.shape[0] - h + 1, X.shape[1] - w + 1))
    for i in range(Y.shape[0]):
        for j in range(Y.shape[1]):
            # 两层循环,相乘,求和
            Y[i, j] = (X[i:i + h, j:j + w] * K).sum()
    # 返回Y
    return Y


# ②实现二维卷积层
class Conv2D(nn.Module):
    def __init__(self, kernel_size):
        super().__init__()
        # 定义权重
        self.weight = nn.Parameter(torch.rand(kernel_size))
        # 定义偏移
        self.bias = nn.Parameter(torch.zeros(1))

    # 定义正向传播
    def forward(self, x):
        return corr2d(x, self.weight) + self.bias

if __name__ == '__main__':
    # 定义模型
    conv2d = nn.Conv2d(1, 1, kernel_size=(1, 2), bias=False)
    # 定义X
    X = torch.ones((6, 8))
    X[:, 2:6] = 0
    # 定义K
    K = torch.tensor([[1.0, -1.0]])
    # 计算Y
    Y = corr2d(X, K)
    X = X.reshape((1, 1, 6, 8))
    Y = Y.reshape((1, 1, 6, 7))
    # 训练10轮
    for i in range(10):
        # 计算Y_hat
        Y_hat = conv2d(X)
        # 损失
        l = (Y_hat - Y)**2
        # 梯度归零
        conv2d.zero_grad()
        # 后向传播
        l.sum().backward()
        # 优化函数 学习率=3e-2
        conv2d.weight.data[:] -= 3e-2 * conv2d.weight.grad
        if (i + 1) % 2 == 0:
            print(f'batch {i+1}, loss {l.sum():.3f}')
    # 经过10轮学习的权重为
    print(conv2d.weight.data.reshape((1, 2)))

结果

python 复制代码
batch 2, loss 1.463
batch 4, loss 0.358
batch 6, loss 0.106
batch 8, loss 0.037
batch 10, loss 0.014
tensor([[ 1.0066, -0.9830]])

Process finished with exit code 0
相关推荐
2601_966949653 分钟前
如何利用 1 分钟 K 线在 14:50 精准捕捉尾盘异动?Python 量化实战:1 分钟 K 线 + 全市场扫描
开发语言·人工智能·python·量化·quantdash·量化数据源
海天一色y4 分钟前
强化学习工具函数详解:从经验回放到优势函数计算
人工智能·python·强化学习
船厂电气自动化ai大模型14 分钟前
AI大模型与数学·第56课 快速傅里叶变换FFT:DFT高效优化算法,图像、音频、扩散模型工程加速核心工具
数据结构·人工智能·深度学习·算法·机器学习
jay神22 分钟前
YOLO还能不能作为模型baseline?
人工智能·深度学习·yolo·cnn·毕业设计
xier_ran35 分钟前
【infra之路】矩阵乘法(GEMM) Tiling 复用机制总结
人工智能·深度学习
阿图灵43 分钟前
OpenCV 算术运算四件套:NOT/AND/OR 位运算与图像混合
图像处理·人工智能·python·opencv·计算机视觉·位运算
边吃番茄边敲代码1 小时前
企业智能助手Agent 安全测试,包含Prompt Injection 和越权工具调用
人工智能·python·功能测试·ai·单元测试·prompt·模块测试
科技小E1 小时前
自建团队、买SaaS、还是用自动化AI算法训练服务器DLTM?企业AI视觉三条路线算笔账
人工智能·深度学习·自动化
阿图灵1 小时前
OpenCV 阈值与模糊:全局/自适应阈值、Canny 边缘检测与三种模糊
图像处理·人工智能·python·opencv·计算机视觉·边缘检测
桃西西呀1 小时前
国内金价破 1000,「涨了多少」和「该买多少」是数学问题
人工智能·python·数据分析