24 for (
int i = 0; i < cols; ++i) {
31 const char *v = getenv(name);
35 if (v[0] ==
'1' || v[0] ==
'y' || v[0] ==
'Y' || v[0] ==
't' || v[0] ==
'T') {
48 static int cached = -1;
65 const int32_t *targets,
70 int force_strict_math)
72 if (!logits || !targets || !d_logits || tokens <= 0 ||
vocab_size <= 0) {
79 const int ignore_index = -100;
80 const int strict = (force_strict_math >= 0) ? (force_strict_math ? 1 : 0)
82 double total_loss = 0.0;
84 int invalid_target_seen = 0;
86 for (
int t = 0; t < tokens; ++t) {
87 const float *row = logits + (size_t)t * (
size_t)
vocab_size;
88 float *drow = d_logits + (size_t)t * (
size_t)
vocab_size;
89 const int target = targets[t];
91 if (target == ignore_index) {
97 invalid_target_seen = 1;
102 double max_logit = (double)row[0];
104 const double rv = (double)row[v];
105 if (rv > max_logit) {
110 double sum_exp = 0.0;
112 sum_exp += exp((
double)row[v] - max_logit);
114 const double inv_sum = 1.0 / sum_exp;
116 drow[v] = (float)(exp((
double)row[v] - max_logit) * inv_sum);
119 total_loss += -(double)row[target] + max_logit + log(sum_exp);
121 float max_logit = row[0];
123 if (row[v] > max_logit) {
128 double sum_exp = 0.0;
130 const float e = expf(row[v] - max_logit);
132 sum_exp += (double)e;
134 const float inv_sum = 1.0f / (float)sum_exp;
139 total_loss += -(double)row[target] + (
double)max_logit + log(sum_exp);
142 drow[target] -= 1.0f;
146 if (valid_tokens > 0) {
147 const float scale = 1.0f / (float)valid_tokens;
148 for (
int t = 0; t < tokens; ++t) {
149 float *drow = d_logits + (size_t)t * (
size_t)
vocab_size;
157 if (invalid_target_seen || valid_tokens == 0) {
160 *loss_out = (float)(total_loss / (
double)valid_tokens);
170 const int32_t *targets,
177 const double scale = 1.0 / (double)tokens;
178 double total_loss = 0.0;
180 for (
int t = 0; t < tokens; ++t) {
181 const float *row = logits + (size_t)t * (
size_t)
vocab_size;
182 float *drow = d_logits + (size_t)t * (
size_t)
vocab_size;
183 const int target = targets[t];
186 double max_logit = (double)row[0];
188 const double rv = (double)row[v];
189 if (rv > max_logit) {
194 double sum_exp = 0.0;
196 sum_exp += exp((
double)row[v] - max_logit);
198 const double inv_sum = 1.0 / sum_exp;
200 const double p = exp((
double)row[v] - max_logit) * inv_sum;
201 drow[v] = (float)(p * scale);
205 const double log_sum_exp = log(sum_exp);
206 const double target_logit = (double)row[target];
207 total_loss += -(target_logit - max_logit - log_sum_exp);
208 drow[target] -= (float)scale;
211 float max_logit = row[0];
213 if (row[v] > max_logit) {
218 double sum_exp = 0.0;
220 const float e = expf(row[v] - max_logit);
222 sum_exp += (double)e;
225 const float inv_sum = 1.0f / (float)sum_exp;
231 const double log_sum_exp = log(sum_exp);
232 const double target_logit = (double)row[target];
233 total_loss += -(target_logit - (double)max_logit - log_sum_exp);
234 drow[target] -= 1.0f;
238 drow[v] *= (float)scale;
244 *loss_out = (float)(total_loss / (
double)tokens);
250 for (
int t = 0; t < tokens; ++t) {
251 const int target = targets[t];
260 const int32_t *targets,
266 if (!logits || !targets || !d_logits || tokens <= 0 ||
vocab_size <= 0) {
275 logits, targets, tokens,
vocab_size, d_logits, loss_out);
280 logits, targets, tokens,
vocab_size, d_logits, loss_out, -1);
293 const int32_t *targets,
305 logits, targets, tokens,
vocab_size, d_logits, loss_out, 1);
int ck_strict_parity_enabled(void)
void softmax_cross_entropy_loss_ptref(const float *logits, const int32_t *targets, int tokens, int vocab_size, float *d_logits, float *loss_out)
static int ce_targets_all_valid_no_ignore(const int32_t *targets, int tokens, int vocab_size)
static int ce_legacy_mode_enabled(void)
static void zero_row_f32(float *row, int cols)
static void softmax_cross_entropy_loss_index_mean_impl(const float *logits, const int32_t *targets, int tokens, int vocab_size, float *d_logits, float *loss_out, int force_strict_math)
static void softmax_cross_entropy_loss_legacy_mean_impl(const float *logits, const int32_t *targets, int tokens, int vocab_size, float *d_logits, float *loss_out)
static int parse_env_bool_on(const char *name)
void softmax_cross_entropy_loss(const float *logits, const int32_t *targets, int tokens, int vocab_size, float *d_logits, float *loss_out)