@Operator(group="nn") public final class ComputeAccidentalHits extends PrimitiveOp
When doing log-odds NCE, the result of this op should be passed through a SparseToDense op, then added to the logits of the sampled candidates. This has the effect of 'removing' the sampled labels that match the true labels by making the classifier sure that they are sampled labels.
Modifier and Type | Class and Description |
---|---|
static class |
ComputeAccidentalHits.Options
Optional attributes for
ComputeAccidentalHits |
operation
Modifier and Type | Method and Description |
---|---|
static ComputeAccidentalHits |
create(Scope scope,
Operand<Long> trueClasses,
Operand<Long> sampledCandidates,
Long numTrue,
ComputeAccidentalHits.Options... options)
Factory method to create a class wrapping a new ComputeAccidentalHits operation.
|
Output<Long> |
ids()
A vector of IDs of positions in sampled_candidates that match a true_label
for the row with the corresponding index in indices.
|
Output<Integer> |
indices()
A vector of indices corresponding to rows of true_candidates.
|
static ComputeAccidentalHits.Options |
seed(Long seed) |
static ComputeAccidentalHits.Options |
seed2(Long seed2) |
Output<Float> |
weights()
A vector of the same length as indices and ids, in which each element
is -FLOAT_MAX.
|
equals, hashCode, op, toString
public static ComputeAccidentalHits create(Scope scope, Operand<Long> trueClasses, Operand<Long> sampledCandidates, Long numTrue, ComputeAccidentalHits.Options... options)
scope
- current scopetrueClasses
- The true_classes output of UnpackSparseLabels.sampledCandidates
- The sampled_candidates output of CandidateSampler.numTrue
- Number of true labels per context.options
- carries optional attributes valuespublic static ComputeAccidentalHits.Options seed(Long seed)
seed
- If either seed or seed2 are set to be non-zero, the random number
generator is seeded by the given seed. Otherwise, it is seeded by a
random seed.public static ComputeAccidentalHits.Options seed2(Long seed2)
seed2
- An second seed to avoid seed collision.public Output<Integer> indices()
public Output<Long> ids()
Copyright © 2022. All rights reserved.