MCPcopy Create free account
hub / github.com/zai-org/CodeGeeX / parallel_lm_logits

Function parallel_lm_logits

codegeex/megatron/model/language_model.py:46–70  ·  view source on GitHub ↗

LM logits using word embedding weights.

(input_, word_embeddings_weight, parallel_output, bias=None)

Source from the content-addressed store, hash-verified

44
45
46def 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
73def get_language_model(

Callers 2

forwardMethod · 0.90
_logits_helperMethod · 0.90

Calls 2

get_argsFunction · 0.90

Tested by

no test coverage detected