| 2256 | } |
| 2257 | |
| 2258 | void ecall_softmax_grad(float *input, float *grad, int N, int M, int C, float *output) { |
| 2259 | float *x = (float*)malloc(sizeof(float)*N); |
| 2260 | memset(x, 0, sizeof(float)*N); |
| 2261 | for(int j=0;j<N;j++){ |
| 2262 | x[j] = 0.0f; |
| 2263 | float sum = input[insert*N+j]; |
| 2264 | float c = 0.0f; |
| 2265 | for(int i=0;i<Nt;i++){ |
| 2266 | float y = input[indexs[i]*N+j] - c; |
| 2267 | float t = sum + y; |
| 2268 | c = (t - sum) - y; |
| 2269 | sum = t; |
| 2270 | } |
| 2271 | x[j] = sum; |
| 2272 | } |
| 2273 | float *g = (float*)malloc(sizeof(float)*N); |
| 2274 | memset(g, 0, sizeof(float)*N); |
| 2275 | for(int j=0;j<N;j++){ |
| 2276 | g[j] = 0.0f; |
| 2277 | float sum = grad[insert*N+j]; |
| 2278 | float c = 0.0f; |
| 2279 | for(int i=0;i<Nt;i++){ |
| 2280 | float y = grad[indexs[i]*N+j] - c; |
| 2281 | float t = sum + y; |
| 2282 | c = (t - sum) - y; |
| 2283 | sum = t; |
| 2284 | } |
| 2285 | g[j] = sum; |
| 2286 | } |
| 2287 | for (int i = 0; i < M; i++) { |
| 2288 | float max_val = x[i * C]; |
| 2289 | for (int j = 1; j < C; j++) { |
| 2290 | max_val = fmaxf(max_val, x[i * C + j]); |
| 2291 | } |
| 2292 | float sum = 0.0f; |
| 2293 | for (int j = 0; j < C; j++) { |
| 2294 | float scaled_input = x[i * C + j] - max_val; |
| 2295 | sum += expf(scaled_input); |
| 2296 | } |
| 2297 | for (int j = 0; j < C; j++) { |
| 2298 | float scaled_input = x[i * C + j] - max_val; |
| 2299 | x[i * C + j] = expf(scaled_input) / sum; |
| 2300 | } |
| 2301 | } |
| 2302 | float *gt = (float *)malloc(sizeof(float) * M); |
| 2303 | memset(gt, 0, sizeof(float) * M); |
| 2304 | for (int i = 0; i < M; i++) { |
| 2305 | gt[i] = 0.0f; |
| 2306 | for (int j = 0; j < C; j++) { |
| 2307 | gt[i] += g[i * C + j] * x[i * C + j]; |
| 2308 | } |
| 2309 | } |
| 2310 | for (int i = 0; i < M; i++) { |
| 2311 | for (int j = 0; j < C; j++) { |
| 2312 | g[i*C+j] = (g[i * C + j] - gt[i]) * x[i * C + j]; |
| 2313 | } |
| 2314 | } |
| 2315 | float mean = 0.0f, std, mean_sqr = 0.0f; |
no test coverage detected