(pfor_input)
| 2591 | |
| 2592 | @RegisterPFor("Multinomial") |
| 2593 | def _convert_multinomial(pfor_input): |
| 2594 | logits, logits_stacked, _ = pfor_input.input(0) |
| 2595 | num_samples = pfor_input.unstacked_input(1) |
| 2596 | seed = pfor_input.get_attr("seed") |
| 2597 | seed2 = pfor_input.get_attr("seed2") |
| 2598 | output_dtype = pfor_input.get_attr("output_dtype") |
| 2599 | logging.warning( |
| 2600 | "Note that Multinomial inside pfor op may not give same output as " |
| 2601 | "inside a sequential loop.") |
| 2602 | |
| 2603 | n = pfor_input.pfor.loop_len_vector[0] |
| 2604 | if logits_stacked: |
| 2605 | flattened_logits = _flatten_first_two_dims(logits) |
| 2606 | samples = gen_random_ops.multinomial( |
| 2607 | flattened_logits, |
| 2608 | num_samples, |
| 2609 | seed=seed, seed2=seed2, output_dtype=output_dtype) |
| 2610 | stacked_samples = _unflatten_first_dim(samples, [n]) |
| 2611 | else: |
| 2612 | samples = gen_random_ops.multinomial( |
| 2613 | logits, num_samples * n, |
| 2614 | seed=seed, seed2=seed2, output_dtype=output_dtype) |
| 2615 | stacked_samples = array_ops.transpose( |
| 2616 | array_ops.reshape(samples, [-1, n, num_samples]), [1, 0, 2]) |
| 2617 | |
| 2618 | return wrap(stacked_samples, True) |
| 2619 | |
| 2620 | |
| 2621 | # linalg_ops |
nothing calls this directly
no test coverage detected