Implement true CLT: per-source-layer encoders/decoders with causal masking#12
Open
bartbussmann wants to merge 1 commit into
Open
Implement true CLT: per-source-layer encoders/decoders with causal masking#12bartbussmann wants to merge 1 commit into
bartbussmann wants to merge 1 commit into
Conversation
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Motivation
num_input_layers) so concatenated multi-layer inputs are handled correctly.Description
num_input_layerstoEncoderConfigand introducedmulti_layer/true_cltflags inSharedEncoderto detect CLT mode vianum_output_layers > 1andnum_input_layers > 1and to computeinput_size_per_layerwhen needed.W_enc/b_encbecome shape(num_input_layers, input_size_per_layer, dict_size)/(num_input_layers, dict_size)in true CLT mode, and addedencode_linearto compute per-layer pre-activations.W_decwith shape(num_input_layers, num_output_layers, dict_size, output_size)in true CLT mode and registered acausal_maskbuffer so source layerscan only write to target layerst >= s, and applied that mask duringdecodeandmake_decoder_weights_and_grad_unit_norm.update_inactive_features, auxiliary-loss path,BatchTopKflattening logic, and forward paths to correctly handle 3D activations ((B, S, D)) produced by CLT encoders and to preserve backward compatibility in non-CLT mode.test_clt_is_causalto verify perturbations to later-layer inputs do not affect earlier-layer outputs.Testing
python -m py_compile base.py config.py test_clt.pywhich reported no syntax errors (success).python test_clt.pybut the run failed in this environment withModuleNotFoundError: No module named 'torch', so the full runtime tests could not be executed here (failure due to missing dependency).Codex Task