| 134 | p_start = model.add_parameters({EMBED_DIM}); |
| 135 | } |
| 136 | Expression loss(ComputationGraph& cg, const Expression& v, const string& code) { |
| 137 | decoder.new_graph(cg); |
| 138 | Expression h = tanh(v); |
| 139 | vector<Expression> init = {v, h}; |
| 140 | decoder.start_new_sequence(init); |
| 141 | Expression start = parameter(cg, p_start); |
| 142 | PrefixNode* cur = &pfc->root; |
| 143 | decoder.add_input(start); |
| 144 | size_t i = 0; |
| 145 | vector<Expression> errs(code.size()); |
| 146 | while(i < code.size()) { |
| 147 | assert(cur); |
| 148 | Expression pred = decoder.back(); |
| 149 | Expression rp = parameter(cg, cur->pred); |
| 150 | Expression bias = parameter(cg, cur->bias); |
| 151 | Expression p = logistic(dot_product(pred, rp) + bias); |
| 152 | // maybe squared error instead of xentropy? |
| 153 | if (code[i] == '0') p = 1.f - p; |
| 154 | errs[i] = log(p); |
| 155 | Expression cond = parameter(cg, code[i] == '0' ? cur->zero_cond : cur->one_cond); |
| 156 | decoder.add_input(cond); |
| 157 | cur = code[i] == '0' ? cur->zero_child : cur->one_child; |
| 158 | ++i; |
| 159 | } |
| 160 | assert(cur->terminal); |
| 161 | return -sum(errs); |
| 162 | } |
| 163 | }; |
| 164 | |
| 165 | template <class Builder> |
no test coverage detected