加载TensorFlow模型时遭遇CORS策略问题,权重获取失败求助

加载模型时出现错误,10小时前突然出现,之前一切正常。尝试编写了以下包装API,目前能获取到模型文件,但权重文件获取失败:
router.get("/modelWrapper", async (req, res) => { try { console.log("********** poseDetector WRAPPER *************"); const response = await fetch('https://www.kaggle.com/models/google/movenet/frameworks/tfJs/variations/singlepose-lightning/versions/4/model.json?tfjs-format=file&tfhub-redirect=true'); console.log("********** RESPONSE WRAPPER *************"); // Check if response is successful if (!response.ok) { throw new Error('Failed to fetch resource'); } // Get response body as JSON const data = await response.json(); // Add CORS headers res.setHeader('Access-Control-Allow-Origin', '*'); console.log("********** HEADER WRAPPER *************"); // Send response back to frontend res.json(data); } catch (error) { console.error('Error:', error.message); res.status(500).send('Internal server error'); } });
问题原因
你只代理了model.json文件,但TensorFlow.js加载模型时,model.json里的权重文件路径是相对Kaggle服务器的地址,前端会直接请求Kaggle的权重文件,依然会碰到CORS或者权限限制,导致权重加载失败。另外10小时前突然出问题,大概率是Kaggle调整了模型的访问策略。
解决办法
1. 代理所有模型相关请求
修改路由,让它能转发模型的所有文件请求(包括权重):
router.get("/modelWrapper/*", async (req, res) => { try { // 提取请求的子路径,拼接成完整的Kaggle模型文件地址 const subPath = req.params[0]; const baseModelUrl = 'https://www.kaggle.com/models/google/movenet/frameworks/tfJs/variations/singlepose-lightning/versions/4'; // 保留原始请求的查询参数 const queryParams = req.url.split('?')[1] ? `?${req.url.split('?')[1]}` : ''; const targetUrl = `${baseModelUrl}/${subPath}${queryParams}`; const response = await fetch(targetUrl); if (!response.ok) { throw new Error('Failed to fetch resource'); } // 转发原始响应的头部信息,确保内容类型正确 const responseHeaders = new Headers(response.headers); responseHeaders.set('Access-Control-Allow-Origin', '*'); const fileBuffer = await response.arrayBuffer(); res.writeHead(response.status, responseHeaders); res.end(Buffer.from(fileBuffer)); } catch (error) { console.error('Error:', error.message); res.status(500).send('Internal server error'); } });
前端加载模型时,使用你的代理地址作为基础路径:
tf.loadLayersModel('/modelWrapper/model.json?tfjs-format=file&tfhub-redirect=true');
2. 本地托管模型文件
直接把Movenet的完整模型包(包括model.json和所有权重文件)下载到本地服务器的静态资源目录,前端直接请求本地文件,彻底绕开第三方的访问限制。
3. 更换模型源地址
改用TensorFlow Hub的官方TFJS模型地址,这类地址通常默认支持CORS:
tf.loadLayersModel('https://tfhub.dev/google/tfjs-model/movenet/singlepose/lightning/4/default/1/model.json');
内容的提问来源于stack exchange,提问作者pranjali chumbhale
相关产品推荐
相关产品推荐

