@@ -880,10 +880,10 @@ def apply_chat_template(self, conversations, **kwargs):
880880 monkeypatch .setattr (
881881 glm_vl_collate , "extract_skipped_token_ids" , lambda processor : torch .empty (0 , dtype = torch .long )
882882 )
883- monkeypatch .setattr (glm_vl_collate , "infer_assistant_mask_boundary_config " , lambda processor : None )
883+ monkeypatch .setattr (glm_vl_collate , "_glm4v_assistant_mask_boundary_config " , lambda processor : None )
884884 monkeypatch .setattr (
885885 glm_vl_collate ,
886- "build_assistant_loss_mask " ,
886+ "_build_glm4v_assistant_loss_mask " ,
887887 lambda example , input_ids , * args , ** kwargs : (input_ids != 0 ).to (dtype = torch .float32 ),
888888 )
889889 examples = [
@@ -906,6 +906,59 @@ def apply_chat_template(self, conversations, **kwargs):
906906 assert processor .tokenizer .padding_side == "left"
907907
908908
909+ @pytest .mark .parametrize (
910+ ("input_ids" , "expected_mask" ),
911+ [
912+ (
913+ [100 , 1 , 102 , 55 , 56 , 15 , 3 , 4 , 100 , 2 , 102 , 55 , 56 , 15 , 5 , 6 ],
914+ [0 , 0 , 0 , 0 , 0 , 0 , 1 , 1 , 0 , 0 , 0 , 0 , 0 , 0 , 1 , 1 ],
915+ ),
916+ (
917+ [100 , 1 , 102 , 55 , 56 , 15 , 3 , 4 , 99 , 99 ],
918+ [0 , 0 , 0 , 0 , 0 , 0 , 1 , 1 , 0 , 0 ],
919+ ),
920+ (
921+ [100 , 1 , 102 , 55 , 56 , 15 , 70 , 71 , 107 , 8 , 102 , 55 , 56 , 15 , 3 , 4 ],
922+ [0 , 0 , 0 , 0 , 0 , 0 , 1 , 1 , 0 , 0 , 0 , 0 , 0 , 0 , 1 , 1 ],
923+ ),
924+ ],
925+ )
926+ def test_glm4v_assistant_mask_uses_role_boundaries_and_virtual_final_terminator (input_ids , expected_mask ):
927+ class _GlmProcessor :
928+ class _Tokenizer :
929+ chat_template = "<|user|>...<|assistant|>...<|observation|>"
930+ eos_token_id = 99
931+
932+ def encode (self , text , add_special_tokens = False ):
933+ return self (text , add_special_tokens = add_special_tokens )["input_ids" ]
934+
935+ def __call__ (self , text , add_special_tokens = False ):
936+ mapping = {
937+ "<|assistant|>\n " : [102 ],
938+ "<|endoftext|>" : [99 ],
939+ "<|system|>\n " : [105 ],
940+ "<|user|>\n " : [100 ],
941+ "<|observation|>\n " : [107 ],
942+ "<think></think>\n " : [55 , 56 , 15 ],
943+ }
944+ return {"input_ids" : mapping [text ]}
945+
946+ tokenizer = _Tokenizer ()
947+
948+ processor = _GlmProcessor ()
949+ boundary_config = glm_vl_collate ._glm4v_assistant_mask_boundary_config (processor )
950+
951+ mask = glm_vl_collate ._build_glm4v_assistant_loss_mask (
952+ {"conversation" : []},
953+ torch .tensor (input_ids ),
954+ processor ,
955+ torch .empty (0 , dtype = torch .long ),
956+ boundary_config ,
957+ )
958+
959+ assert mask .tolist () == expected_mask
960+
961+
909962def test_expand_image_tokens_handles_multiple_images_and_temporal_grids ():
910963 image_token_id = 163605
911964 input_ids = torch .tensor ([11 , image_token_id , 22 , image_token_id , 33 ])
0 commit comments