""" Patch public MMDuet2 proactive_eval/inference.py for Lenormand HumOmni Track 2. This script mirrors the runtime patch applied in the final notebook. It is intentionally a patch script rather than a copied full inference.py because the base MMDuet2 repository is public and may be cloned by reviewers. """ from pathlib import Path import re import argparse def patch_inference(path: str): inf = Path(path) assert inf.exists(), f"inference.py not found: {inf}" src = inf.read_text(encoding="utf-8") log = [] if 'try:\n from torchvision.io import read_video' not in src: src = src.replace( 'from torchvision.io import read_video', 'try:\n from torchvision.io import read_video\nexcept Exception:\n read_video=None' ) log.append('read_video_safe_import') if 'max_new_tokens=512,' in src: src = src.replace( 'max_new_tokens=512,', 'max_new_tokens=64,\n repetition_penalty=1.3,\n no_repeat_ngram_size=3,' ) log.append('short_generation_and_repetition_control') if " top_k: int = 40\n writeback_mode" not in src: src = src.replace( ' top_k: int = 40\n', " top_k: int = 40\n writeback_mode: str = 'full'\n", 1 ) log.append('add_writeback_mode_arg') if 'self.writeback_mode = getattr' not in src: src = src.replace( ' self.system_prompt = args.system_prompt', ' self.system_prompt = args.system_prompt\n self.writeback_mode = getattr(args, "writeback_mode", "full")\n self.submission_history = []', 1 ) log.append('init_writeback_state') if 'self.submission_history = list()' not in src: src = src.replace( ' self.history = list()', ' self.history = list()\n self.submission_history = list()' ) log.append('reset_submission_history') old = '\n'.join([ ' self.past_key_values = model_output.past_key_values', ' output_token_ids = model_output.sequences', ' output_token_ids = output_token_ids[:, inputs.input_ids.size(1):]', ' reply_text = self.processor.batch_decode(output_token_ids, skip_special_tokens=True)[0]', " if query.get('must_reply', False):", ' reply_text = self.must_reply_prompt + reply_text', " self.history.append({'role': 'assistant', 'content': reply_text, 'time': self.video_time})" ]) new = '\n'.join([ ' self.past_key_values = model_output.past_key_values', ' output_token_ids = model_output.sequences', ' output_token_ids = output_token_ids[:, inputs.input_ids.size(1):]', ' reply_text = self.processor.batch_decode(output_token_ids, skip_special_tokens=True)[0]', " if query.get('must_reply', False):", ' reply_text = self.must_reply_prompt + reply_text', " self.submission_history.append({'role':'assistant','content':reply_text,'time':self.video_time})", " _rt = ' '.join(reply_text.strip().split())", " _is_noreply = _rt.upper() in ('NO REPLY','NO REPLAY','')", " if self.writeback_mode == 'full':", " self.history.append({'role':'assistant','content':reply_text,'full_content':reply_text,'time':self.video_time})", ' else:', " if _is_noreply: short_text='NO REPLY'", " elif self.writeback_mode=='minimal': short_text='ANSWERED.'", " elif self.writeback_mode=='summary': short_text='ANSWERED: '+reply_text.strip().split('.')[0][:60]+'.'", " else: short_text='ANSWERED.'", ' target_len = self.past_key_values.get_seq_length() - output_token_ids.size(1)', ' try:', ' self.past_key_values.crop(target_len)', ' except Exception as e:', " print('crop fallback full:', e)", " self.history.append({'role':'assistant','content':reply_text,'full_content':reply_text,'time':self.video_time})", " if debug_print: print('kvcache length now:', self.past_key_values.get_seq_length())", ' return', " self.history.append({'role':'assistant','content':short_text,'full_content':reply_text,'time':self.video_time})", " short_wrapped = '<|im_start|>assistant\\n' + short_text + '<|im_end|>\\n'", " short_ids = self.processor.tokenizer(short_wrapped, return_tensors='pt', add_special_tokens=False).input_ids.to('cuda:0')", ' with torch.no_grad():', ' _ = self.model(input_ids=short_ids, past_key_values=self.past_key_values, use_cache=True)', ' _clen = self.past_key_values.get_seq_length()', " self.model.model.all_keep_masks = [torch.ones((1, _clen), dtype=torch.bool, device='cuda:0')]" ]) if 'submission_history.append' not in src and old in src: src = src.replace(old, new) log.append('summary_minimal_writeback_keepmask') if "turn['content'] = turn['full_content']" not in src: pp = '\n'.join([ 'def post_process_conversation_for_print(conversation):', ' new_conversation = list()', ' for turn in conversation:', " if isinstance(turn['content'], list):", " res = ''", " for content in turn['content']:", " if 'text' in content:", " res += content['text'].strip()", " turn['content'] = res", " if turn['role'] == 'assistant':", " if 'full_content' in turn:", " turn['content'] = turn['full_content']", " _c = ' '.join(str(turn['content']).strip().split()).upper()", " if _c not in ('NO REPLY', 'NO REPLAY', ''):", ' new_conversation.append(turn)', " elif turn['role'] == 'user':", " if turn['content']:", ' new_conversation.append(turn)', ' return new_conversation\n' ]) src, n = re.subn(r'def post_process_conversation_for_print\(conversation\):.*?return new_conversation\n', pp, src, flags=re.DOTALL) if n == 1: log.append('postprocess_full_content_and_drop_noreply') inf.write_text(src, encoding="utf-8") v = inf.read_text(encoding="utf-8") checks = { 'short_generation': 'max_new_tokens=64' in v, 'writeback_arg': 'writeback_mode: str' in v, 'writeback_history': 'submission_history.append' in v, 'keepmask': 'all_keep_masks = [torch.ones' in v, 'full_content_output': "turn['content'] = turn['full_content']" in v, } print('patch log:', log) print('checks:', checks) assert all(checks.values()), 'patch incomplete' if __name__ == '__main__': parser = argparse.ArgumentParser() parser.add_argument('inference_py', help='Path to public MMDuet2 proactive_eval/inference.py') args = parser.parse_args() patch_inference(args.inference_py)