(logit, classes)
| 154 | |
| 155 | |
| 156 | def classify_step(logit, classes): |
| 157 | class_vector = F.softmax(logit, 1).data.squeeze() |
| 158 | assert len(class_vector) == len(classes), "class number must match" |
| 159 | probs, idx = class_vector.sort(0, True) |
| 160 | result = classes[idx[0]] |
| 161 | return result |