You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

RNN采用显式梯度在CUDA运行报错,CPU环境正常求助

问题:Flux.jl RNN在GPU上显式梯度计算触发非法内存访问错误

我用Flux.jl实现的RNN在CPU上运行正常,但切换到GPU(显式梯度计算)时触发ERROR_ILLEGAL_ADDRESS(代码700),错误栈如下:

ERROR: CUDA error: an illegal memory access was encountered (code 700, ERROR_ILLEGAL_ADDRESS)
Stacktrace:
 [1] throw_api_error(res::CUDA.cudaError_enum)
   @ CUDA C:\Users\XXX\.julia\packages\CUDA\nIZkq\lib\cudadrv\libcuda.jl:27
 [2] isdone
   @ C:\Users\XXX\.julia\packages\CUDA\nIZkq\lib\cudadrv\stream.jl:111 [inlined]
 [3] spinning_synchronization(f::typeof(CUDA.isdone), obj::CuStream)
   @ CUDA C:\Users\XXX\.julia\packages\CUDA\nIZkq\lib\cudadrv\synchronization.jl:79
 [4] device_synchronize(; blocking::Bool, spin::Bool)
   @ CUDA C:\Users\XXX\.julia\packages\CUDA\nIZkq\lib\cudadrv\synchronization.jl:171
 [5] device_synchronize()
   @ CUDA C:\Users\XXX\.julia\packages\CUDA\nIZkq\lib\cudadrv\synchronization.jl:169
 [6] top-level scope
   @ C:\Users\XXX\.julia\packages\CUDA\nIZkq\src\initialization.jl:210

caused by: CUDA error: an illegal memory access was encountered (code 700, ERROR_ILLEGAL_ADDRESS)
Stacktrace:
  [1] throw_api_error(res::CUDA.cudaError_enum)
    @ CUDA C:\Users\XXX\.julia\packages\CUDA\nIZkq\lib\cudadrv\libcuda.jl:27

为支持时间反向传播(BPTT),我已按Flux文档要求将输入设为时间步向量的特征向量,简化后的复现代码如下(移除训练循环):

using Flux
using ChainRulesCore
using CUDA

dev=gpu # cpu is working fine

m = Chain(RNN(2 => 5), Dense(5 => 1)) |> dev

x = [rand(Float32, 2) for i = 1:3] |> dev;
y = [rand(Float32, 1) for i=1:1] |> dev

[m(xi) for xi in x]

using Flux.Losses: mse

function loss(m, x, y)
    @ignore_derivatives Flux.reset!(m)
    m(x[1]) # ignores the output but updates the hidden states
    m(x[2]) # ignore second output
    mse(m(x[3]),y[1])
end
  
loss(m, x, y)

grads = Flux.gradient(m, x, y) do m,x,y
    loss(m, x, y)
end

optim = Flux.setup(Flux.Adam(), m) 
Flux.update!(optim, m, grads[1])

当前环境版本:Julia 1.9.3、CUDA v5.1.0、ChainRulesCore v1.18.0、Flux v0.14.6。想确认新版本Flux中显式梯度的RNN是否完全支持CUDA?


分析与解决方案
  • 版本兼容性修复:你使用的Flux v0.14.6属于旧版本,后续Flux v0.15+对CUDA环境下RNN的梯度计算逻辑做了大量修复,显式梯度模式的GPU支持已完善。建议升级Flux到最新稳定版,同时同步更新CUDA.jl到匹配版本(CUDA.jl v5.x兼容Julia 1.9及最新Flux)。
  • 代码写法优化:当前手动分步调用RNN的方式在GPU上易引发状态同步问题,可改为带状态的RNN调用模式,避免直接修改模型内部状态:
    function loss(m, x, y)
        state = Flux.reset!(m)
        state, _ = m(state, x[1])
        state, _ = m(state, x[2])
        state, out = m(state, x[3])
        mse(out, y[1])
    end
    
    也可将输入转为批量时间步格式(如Float32[2,3]数组,时间步为第二维度),用Flux.batchseq处理,更适配GPU计算模式。
  • 错误根源:旧版本Flux中,RNN隐藏状态在GPU梯度回传时存在内存越界风险,尤其是显式梯度模式下的状态跟踪逻辑不完善,升级版本后这类问题已被修复。
  • 新版本支持情况:Flux v0.15及以上版本完全支持CUDA环境下的显式梯度RNN计算,包括BPTT场景,社区已验证大量同类案例。

内容的提问来源于stack exchange,提问作者abj

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.06 16:33:19