LM logits using word embedding weights.
(input_, word_embeddings_weight, parallel_output, bias=None)
| 44 | |
| 45 | |
| 46 | def parallel_lm_logits(input_, word_embeddings_weight, parallel_output, bias=None): |
| 47 | """LM logits using word embedding weights.""" |
| 48 | # Parallel logits. |
| 49 | input_parallel = mpu.copy_to_tensor_model_parallel_region(input_) |
| 50 | # Matrix multiply. |
| 51 | args = get_args() |
| 52 | if args.shrink_logit_embedding_gradient: |
| 53 | if hasattr(args, 'iteration'): |
| 54 | alpha = get_shrink_embedding_gradient_alpha(args.iteration + 1) |
| 55 | else: |
| 56 | alpha = args.shrink_embedding_gradient_alpha |
| 57 | word_embeddings_weight = word_embeddings_weight if alpha == 1.0 \ |
| 58 | else ( |
| 59 | word_embeddings_weight * alpha + |
| 60 | word_embeddings_weight.detach() * (1 - alpha) |
| 61 | ) |
| 62 | if bias is None: |
| 63 | logits_parallel = F.linear(input_parallel, word_embeddings_weight.half()) |
| 64 | else: |
| 65 | logits_parallel = F.linear(input_parallel, word_embeddings_weight.half(), bias) |
| 66 | # Gather if needed. |
| 67 | if parallel_output: |
| 68 | return logits_parallel |
| 69 | |
| 70 | return mpu.gather_from_tensor_model_parallel_region(logits_parallel) |
| 71 | |
| 72 | |
| 73 | def get_language_model( |
no test coverage detected