chenbhao commited on
Commit
5a72179
·
1 Parent(s): 2dd2738

Revert camera prompt changes; add NaN guard and CUDA sync; add engineering challenges doc

Browse files

- prepare_turn: restored original '请描述这张图片' prompt + image tokens after text
- build_ids: removed system prompt and image_context parameter
- model.py stream_generate: added nan_to_num guard on logits/probs
- omni_o_call.py: added torch.cuda.synchronize() after ASR before generation
- docs/interview/challenges.md: new document covering real engineering challenges
- docs/interview/index.md: added link to challenges.md

docs/interview/challenges.md ADDED
@@ -0,0 +1,292 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # 面试:工程挑战与解决方案
2
+
3
+ > omni-o 实时语音通话集成中遇到的实际工程问题,覆写/转换/部署/数值稳定性等方面
4
+
5
+ ## Q1. 原始 .pth 权重 → HF safe tensors 格式转换
6
+
7
+ ### 问题
8
+
9
+ omni-o 的原始发布权重为原生的 `.pth` 文件,但 HuggingFace 生态需要 `config.json` + `model.safetensors` + `modeling_xxxx.py` 结构。需要自动检测并兼容两种加载方式。
10
+
11
+ ### 解决方案
12
+
13
+ 在 `omni_o_call.py:307` 实现自动检测:
14
+
15
+ ```python
16
+ is_hf = os.path.exists(os.path.join(ckpt_dir, 'config.json')) and \
17
+ (os.path.exists(os.path.join(ckpt_dir, 'model.safetensors')) or
18
+ os.path.exists(os.path.join(ckpt_dir, 'pytorch_model.bin')))
19
+
20
+ if is_hf:
21
+ model = VAM.from_pretrained(ckpt_dir, ...)
22
+ else:
23
+ # 原始 .pth 路径
24
+ state = torch.load(ckpt_path, ...)
25
+ model.load_state_dict(state, strict=False)
26
+ ```
27
+
28
+ 转换脚本 `scripts/convert_omni_o_to_hf.py` 负责:
29
+ 1. 加载原始 `.pth` checkpoint
30
+ 2. 构建 `VAMConfig` + `VAM` 实例
31
+ 3. 用 `load_state_dict` 注入权重
32
+ 4. 通过 `save_pretrained` 输出 `config.json` + `model.safetensors`
33
+ 5. 从 `checkpoint/omni/native_hf` 复制 tokenizer 文件
34
+ 6. 将 `modeling_omni_o.py` 写入目标目录,实现 `trust_remote_code`
35
+
36
+ ### 遇到的坑
37
+
38
+ - **Meta init 冲突**:`from_pretrained` 内部会先以 meta device 初始化,但 audio/vision encoder 在 meta 下无法加载。解决:重写 `VAM.from_pretrained()` 方法,在 meta init 阶段跳过外部 encoder。
39
+
40
+ ```python
41
+ @classmethod
42
+ def from_pretrained(cls, path, audio_encoder_path=None, vision_model_path=None, **kwargs):
43
+ kwargs['config'] = VAMConfig.from_pretrained(path)
44
+ with contextlib.redirect_stdout(io.StringIO()):
45
+ model = super().from_pretrained(path, torch_dtype=torch.float16, **kwargs)
46
+ # 外部 encoder 在 meta init 之后单独加载
47
+ if audio_encoder_path:
48
+ model.audio_encoder = cls._init_audio_encoder(audio_encoder_path)[0]
49
+ if vision_model_path:
50
+ model.vision_encoder, model.vision_processor = cls._init_vision_encoder(vision_model_path)
51
+ return model
52
+ ```
53
+
54
+ - **权重不匹配**:原始 checkpoint 的 key 命名与 HF 规范不同,需要 `strict=False` + 手动处理缺失/多余 key。
55
+
56
+ ---
57
+
58
+ ## Q2. FP16 数值溢出导致 CUDA 崩溃
59
+
60
+ ### 现象
61
+
62
+ 模型加载为 `model.half()` 后,无摄像头时推理正常;打开摄像头后,`multinomial` 调用随机崩溃:
63
+
64
+ ```
65
+ Error generating frame: CUDA error: device-side assert triggered
66
+ probability tensor contains inf, nan, or negative
67
+ ```
68
+
69
+ ### 根因分析
70
+
71
+ 调试过程:
72
+
73
+ 1. **定位崩溃点**:在 `stream_generate:369` 的 `torch.multinomial(F.softmax(logits), 1)`
74
+ 2. **检查 logits**:打印 `logits` 统计 — 正常情况下无 NaN,但特定图像输入时出现
75
+ 3. **追溯 NaN 起源**:检查 forward pass 输出 `out.logits` — 只在 `pixel_values` 非空时出现
76
+ 4. **量化分析**:测量 vision projector 输出的统计量
77
+
78
+ ```python
79
+ vision_tensors.min() = -25.6, .max() = 26.6, .std() = 4.7
80
+ ```
81
+
82
+ 5. **理论推算**:Transformer attention 中 `Q @ K^T / sqrt(d)` 可能超出 FP16 范围
83
+
84
+ 分析详细计算:
85
+ - Hidden state 来自 embedding,值约 ±0.036
86
+ - Vision feature 替换后,值约 ±26(差异 ~720 倍)
87
+ - 注意力分数 `Q @ K^T / sqrt(96)` 可能在 `±26 × ±26 × 96 / 9.8 ≈ 6600` — 接近但未立即溢出
88
+ - 但连续多层 Transformer 的中间激活可能累积放大
89
+
90
+ 实测发现:用 syntheitc 图(纯色、随机噪点)均不崩溃,只有真实摄像头画面会触发出问题。说明像素分布差异导致 SigLIP 输出极端值。
91
+
92
+ ### 尝试过的方案与权衡
93
+
94
+ | 方案 | 效果 | 问题 |
95
+ | --- | --- | --- |
96
+ | `vision_tensors.clamp(-10, 10)` | 防止溢出,正常测试通过 | 扭曲了 60% 的特征值(±25→±10),信息损失严重 |
97
+ | `vision_tensors.clamp(-30, 30)` | 几乎无扭曲 | 限幅太宽,无法阻止溢出 |
98
+ | Projector 输出 RMS norm 后匹配 hidden_states scale | 理论合理 | 改变特征分布,可能影响生成质量 |
99
+ | Thinker 转为 bfloat16 | 保留 fp32 动态范围,无溢出 | 需验证 GPU 兼容性(检查 `torch.cuda.is_bf16_supported()`) |
100
+ | Thinker 转为 float32 | 彻底解决 | 参数量翻倍(63M→252MB,实际可接受) |
101
+
102
+ ### 最终方案
103
+
104
+ 在 `stream_generate` 中加入 NaN guard,用 `torch.nan_to_num` 兜底:
105
+
106
+ ```python
107
+ logits = out.logits[0, -1, :].clone().float() / (temperature + 1e-9)
108
+ logits = torch.nan_to_num(logits, nan=-100.0, posinf=-100.0, neginf=-100.0)
109
+ probs = F.softmax(logits, dim=-1)
110
+ probs = torch.nan_to_num(probs)
111
+ if probs.sum() <= 0:
112
+ probs = torch.ones_like(probs) / probs.shape[-1]
113
+ text_token = torch.multinomial(probs, 1).item()
114
+ ```
115
+
116
+ 这是「救火」方案而不是根除方案。根因是 FP16 动态范围不够,彻底解决需要将 thinker 层转为 bfloat16 或 float32。
117
+
118
+ ### 面试价值
119
+
120
+ 这个问题展现了:
121
+ 1. **调试方法论**:从 crash �� → logits → forward pass → 逐层追溯的思维链条
122
+ 2. **数值分析能力**:能手动估算 FP16 溢出边界,理解浮点数表示
123
+ 3. **工程取舍**:理解 clamp 的信息损失,选择兜底而非限幅
124
+ 4. **对精度的理解**:FP16 vs BF16 vs FP32 的动态范围设计差异
125
+
126
+ ---
127
+
128
+ ## Q3. CUDA 多模型并发导致设备端断言
129
+
130
+ ### 现象
131
+
132
+ ASR(funasr)与主模型(VAM)同时使用同一 GPU,在 ASR 完成后立即启动生成时偶发崩溃。
133
+
134
+ ### 根因
135
+
136
+ `prepare_turn` 中 `asr_run(samples)` 在 `MODEL_LOCK` 外执行。funasr 使用自定义 CUDA stream,其 kernel launch 是异步的。当 `asr_run` 返回、主模型立即在默认 stream 上启动 forward 时,两个 stream 上的操作可能乱序执行,导致:
137
+
138
+ - 内存竞争:funasr 释放的显存被主模型复用,但仍有未完成的 kernel 在读取
139
+ - CUDA 设备端 assert:触发非法内存访问
140
+
141
+ ### 解决方案
142
+
143
+ 在 ASR 后、`run_generate` 前插入 CUDA 同步:
144
+
145
+ ```python
146
+ if torch.cuda.is_available():
147
+ torch.cuda.synchronize()
148
+ ```
149
+
150
+ 同时在 SSE 路径和 WebSocket 路径均加入此保护。
151
+
152
+ ### 更根本的修复
153
+
154
+ 将 ASR 也纳入 `MODEL_LOCK` 保护,或为 ASR 使用独立 CUDA stream 并显式同步。当前 `synchronize()` 是轻量级修复。
155
+
156
+ ---
157
+
158
+ ## Q4. 训练偏见 vs 提示工程的矛盾
159
+
160
+ ### 现象
161
+
162
+ 无论系统提示如何写("不要描述画面"、"用视觉作为语境"),模型仍然会先详细描述摄像头画面内容,再回应问题。
163
+
164
+ ### 根因
165
+
166
+ omni-o 的训练数据中,`<|image_pad|>` token 出现的位置与「请描述这张图片」紧密关联。模型在训练阶段学到的是:看见图像 → 先描述。这种权重级别的关联无法通过 prompt engineering 消除。
167
+
168
+ ### 尝试过的方案
169
+
170
+ | 方案 | 效果 |
171
+ | --- | --- |
172
+ | 移除 prompt 中"请描述这张图片" | 几乎无改善 |
173
+ | 图像 token 放在用户文字之前 | 模型更早看到图像 → 更早开始描述 |
174
+ | 加上 `[不描述画面]` 前缀 | 无明显作用 |
175
+ | 加 system prompt "Do not describe unless asked" | 几乎无改善 |
176
+ | 图像 token 放在用户文字之后 | 模型描述完才看文字 → 更差 |
177
+
178
+ ### 正确的解决方向
179
+
180
+ 需要 SFT(Supervised Fine-Tuning)数据,其中图像 token 后跟着自然对话(而非描述),让模型学习到:图像 token 可以用于回答与图像相关的问题,而不仅仅触发描述行为。
181
+
182
+ 具体做法:
183
+ 1. 构造多轮对话 SFT 数据:用户提问 → 模型用视觉信息回答(不描述)
184
+ 2. 配对的图像+问答对(例如 VQA 数据集改造)
185
+ 3. 冻结 vision encoder,只训练 projector + thinker adapter
186
+
187
+ ---
188
+
189
+ ## Q5. HF `from_pretrained` 的 meta init 冲突
190
+
191
+ ### 问题
192
+
193
+ `PreTrainedModel.from_pretrained()` 内部会在 `meta` device 上创建模型,再加载权重。但 VAM 的 `__init__` 中会调用 encoder 初始化函数,这些函数无法在 meta device 上运行(需要下载模型、加载权重等)。
194
+
195
+ ### 解决方案
196
+
197
+ 重写 `from_pretrained` 类方法,绕过 `meta` 设备上的 encoder 初始化:
198
+
199
+ ```python
200
+ @classmethod
201
+ def from_pretrained(cls, pretrained_model_name_or_path, *model_args,
202
+ audio_encoder_path=None, vision_model_path=None, **kwargs):
203
+ config = VAMConfig.from_pretrained(pretrained_model_name_or_path, **kwargs)
204
+ # 用 HF 原生方法加载,但准备好在 meta init 后注入 encoder
205
+ model = super().from_pretrained(pretrained_model_name_or_path, *model_args, **kwargs)
206
+ model.audio_encoder = cls._init_audio_encoder(audio_encoder_path) if audio_encoder_path else None
207
+ model.vision_encoder, model.vision_processor = (
208
+ cls._init_vision_encoder(vision_model_path) if vision_model_path else (None, None))
209
+ return model
210
+ ```
211
+
212
+ 核心技巧:先让 HF 处理内部模块(thinker,talker,projectors),再手动初始化外部 encoder。因为 encoder 是 `nn.Module` 属性而非 `PreTrainedModel` 的子模块,HF 不会自动管理它们。
213
+
214
+ ---
215
+
216
+ ## Q6. VAD + ASR + 生成的多线程实时架构
217
+
218
+ ### 架构设计
219
+
220
+ ```
221
+ WebSocket 接收音频
222
+
223
+
224
+ SileroVAD(ONNX Runtime,CPU)
225
+ │ 检测到语音结束
226
+
227
+ 音频缓冲 → ASR(funasr, CUDA)
228
+
229
+
230
+ 摄像头画面 + ASR 文字 → prep_image + build_ids
231
+
232
+
233
+ MODEL_LOCK → VAM.generate() → 流式文本 + 音频 code
234
+
235
+
236
+ MimiDecoder → PCM → WebSocket 返回
237
+ ```
238
+
239
+ ### 关键工程细节
240
+
241
+ 1. **VAD interrupt**: 生成过程中,后台线程持续接收 WebSocket 音频并送入 VAD。如果检测到新语音,设置 `session.interrupt = True`,生成循环在每步检查并中断:
242
+
243
+ ```python
244
+ for y, af in run_generate(x, audio_inputs, audio_lens, pixel_values, ...):
245
+ if poll_interrupt() or session.interrupt:
246
+ interrupted = True
247
+ break
248
+ ```
249
+
250
+ 2. **线程安全**:多个 threading 模块使用 `MODEL_LOCK` 串行化生成;`inference_mode` 确保不存梯度;每个 session 有自己的 `RealtimeSession` 实例管理 VAD 状态。
251
+
252
+ 3. **音频流**:生成时 8 层音频 code 交错输出,通过 `stream_pcm` 逐步解码为 PCM 并以 base64 分片推送。
253
+
254
+ ### 面试价值
255
+
256
+ 展示了对实时系统设计的理解:低延迟音频处理、中断机制、线程安全、资源锁。
257
+
258
+ ---
259
+
260
+ ## Q7. 音频 code 的流式解码与重叠播放
261
+
262
+ ### 问题
263
+
264
+ 音频 codec(Mimi)一次解码一帧,但生成时 8 层 code 是逐 token 产生的。需要在生成完一个「音频帧」(8 个 code,每个层一个)后立即解码并播放,同时后续帧在生成中。
265
+
266
+ ### 实现
267
+
268
+ ```python
269
+ def stream_pcm(frames, flush=False):
270
+ cf, ov_max = cfg.audio_chunk_frames, cfg.audio_overlap
271
+ if not flush and n >= cf and n % cf == 0:
272
+ ov = min(ov_max, n - cf)
273
+ p = pcm_bytes(frames[-(cf + ov):], ov)
274
+ if p: yield p
275
+ ```
276
+
277
+ 关键参数:
278
+ - `audio_chunk_frames=4`:每 4 帧解码一次
279
+ - `audio_overlap=2`:保留 2 帧重叠,避免帧边界 click 噪音
280
+ - 重叠部分:`pcm_bytes` 中根据帧时长计算切除的样本数
281
+
282
+ ---
283
+
284
+ ## 总结:面试中可以讲的故事主线
285
+
286
+ 1. **「有个模型从原始 .pth 转成 HF 格式」** → meta init 冲突 → 重写 `from_pretrained`
287
+ 2. **「用户一开摄像头就崩溃」** → 逐层追查到 FP16 溢出 → NaN guard + 数值分析
288
+ 3. **「ASR 和模型打架」** → CUDA stream 乱序 → `synchronize`
289
+ 4. **「模型永远先描述画面再回答」** → 训练数据偏见无法用 prompt 修复 → 需要 SFT
290
+ 5. **「实时语音要 200ms 内响应」** → VAD interrupt + 流式解码 + thread safety
291
+
292
+ 每个故事都展示了:发现问题 → 分析根因 → 理解取舍 → 实施修复 → 思考更优方案 的完整工程思维链条。
docs/interview/index.md CHANGED
@@ -11,6 +11,7 @@
11
  | [多模态](multimodal.md) | VLM/VAM 注入范式、SigLIP/SenseVoice、TalkerModule、流式生成、冻结策略 |
12
  | [MoE](moe.md) | 路由机制、负载均衡、死 Expert 梯度、Expert 初始化、推理优化 |
13
  | [推理优化](inference.md) | KV Cache、Prefill/Decode、采样策略、Repetition Penalty、量化、显存估算 |
 
14
 
15
  ---
16
 
 
11
  | [多模态](multimodal.md) | VLM/VAM 注入范式、SigLIP/SenseVoice、TalkerModule、流式生成、冻结策略 |
12
  | [MoE](moe.md) | 路由机制、负载均衡、死 Expert 梯度、Expert 初始化、推理优化 |
13
  | [推理优化](inference.md) | KV Cache、Prefill/Decode、采样策略、Repetition Penalty、量化、显存估算 |
14
+ | [工程挑战](challenges.md) | HF 转换、FP16 溢出、CUDA 同步、训练偏见、实时架构、流式音频 |
15
 
16
  ---
17
 
scripts/omni_o_call.py CHANGED
@@ -33,16 +33,11 @@ def prep_image(b64):
33
  img = Image.open(io.BytesIO(base64.b64decode(b64))).convert('RGB')
34
  return {k: v.to(M['device']) for k, v in M['model'].vision_processor(images=img, return_tensors="pt").items()}
35
 
36
- def build_ids(prompt, history, image_context=False):
37
  tok, dev = M['tokenizer'], M['device']
38
  cfg = M['cfg']
39
  hist = history[-cfg.max_history_turns:] if cfg.max_history_turns > 0 else []
40
- if image_context:
41
- sys_prompt = "You have live camera input. Use visual information as conversation context naturally. Do not describe what you see unless asked."
42
- msgs = [{"role": "system", "content": sys_prompt}]
43
- else:
44
- msgs = []
45
- msgs += hist + [{"role": "user", "content": prompt}]
46
  t = tok.apply_chat_template(msgs, tokenize=False, add_generation_prompt=True)
47
  return torch.tensor(tok(t)['input_ids'], dtype=torch.long, device=dev)[None, ...]
48
 
@@ -117,11 +112,7 @@ def prepare_turn(text, samples, image_b64, do_asr_for_image):
117
  if image_b64:
118
  pixel_values = prep_image(image_b64)
119
  m = M['model']
120
- img_tokens = m.config.image_special_token * m.config.image_token_len
121
- if prompt.strip():
122
- prompt = "[不描述画面]\n" + img_tokens + "\n" + prompt
123
- else:
124
- prompt = img_tokens
125
  return audio_inputs, audio_lens, pixel_values, prompt, user_text, asr_thread, asr_result
126
 
127
 
@@ -161,7 +152,9 @@ def init_web_app():
161
  def gen():
162
  audio_inputs, audio_lens, pixel_values, prompt, user_text, asr_th, asr_res = prepare_turn(
163
  d.get('text', ''), samples, d.get('image'), do_asr_for_image=True)
164
- x = build_ids(prompt, history, image_context=d.get('image') is not None)
 
 
165
  asr_sent = False
166
  if user_text and samples is not None and d.get('image'):
167
  yield sse({'type': 'user_prompt', 'content': user_text}); asr_sent = True
@@ -260,11 +253,12 @@ def init_web_app():
260
  session.generating = True
261
  audio = session.get_audio()
262
  ws.send(json.dumps({'type': 'generating'}))
263
- has_image = state.get('image') is not None
264
  audio_inputs, audio_lens, pixel_values, prompt, user_text, asr_th, asr_res = prepare_turn(
265
  '', audio, state['image'], do_asr_for_image=True)
266
  if state['image']: state['image'] = None
267
- x = build_ids(prompt, state['history'], image_context=has_image)
 
 
268
  va_rt = voice_args(state.get('voice', 'default'))
269
 
270
  frames, full_text, interrupted = [], '', False
 
33
  img = Image.open(io.BytesIO(base64.b64decode(b64))).convert('RGB')
34
  return {k: v.to(M['device']) for k, v in M['model'].vision_processor(images=img, return_tensors="pt").items()}
35
 
36
+ def build_ids(prompt, history):
37
  tok, dev = M['tokenizer'], M['device']
38
  cfg = M['cfg']
39
  hist = history[-cfg.max_history_turns:] if cfg.max_history_turns > 0 else []
40
+ msgs = hist + [{"role": "user", "content": prompt}]
 
 
 
 
 
41
  t = tok.apply_chat_template(msgs, tokenize=False, add_generation_prompt=True)
42
  return torch.tensor(tok(t)['input_ids'], dtype=torch.long, device=dev)[None, ...]
43
 
 
112
  if image_b64:
113
  pixel_values = prep_image(image_b64)
114
  m = M['model']
115
+ prompt = (prompt + "\n\n" if prompt else "") + "请描述这张图片\n\n" + m.config.image_special_token * m.config.image_token_len
 
 
 
 
116
  return audio_inputs, audio_lens, pixel_values, prompt, user_text, asr_thread, asr_result
117
 
118
 
 
152
  def gen():
153
  audio_inputs, audio_lens, pixel_values, prompt, user_text, asr_th, asr_res = prepare_turn(
154
  d.get('text', ''), samples, d.get('image'), do_asr_for_image=True)
155
+ if torch.cuda.is_available():
156
+ torch.cuda.synchronize()
157
+ x = build_ids(prompt, history)
158
  asr_sent = False
159
  if user_text and samples is not None and d.get('image'):
160
  yield sse({'type': 'user_prompt', 'content': user_text}); asr_sent = True
 
253
  session.generating = True
254
  audio = session.get_audio()
255
  ws.send(json.dumps({'type': 'generating'}))
 
256
  audio_inputs, audio_lens, pixel_values, prompt, user_text, asr_th, asr_res = prepare_turn(
257
  '', audio, state['image'], do_asr_for_image=True)
258
  if state['image']: state['image'] = None
259
+ if torch.cuda.is_available():
260
+ torch.cuda.synchronize()
261
+ x = build_ids(prompt, state['history'])
262
  va_rt = voice_args(state.get('voice', 'default'))
263
 
264
  frames, full_text, interrupted = [], '', False
src/models/vam/model.py CHANGED
@@ -357,6 +357,7 @@ class VAM(LMForCausalLM):
357
  past_kvs = out.past_key_values
358
 
359
  logits = out.logits[0, -1, :].clone().float() / (temperature + 1e-9)
 
360
  if rp != 1.0:
361
  seen = list(set(input_ids[0].tolist()))
362
  score = logits[seen]
@@ -366,7 +367,11 @@ class VAM(LMForCausalLM):
366
  mask = torch.cumsum(F.softmax(sorted_l, dim=-1), dim=-1) > top_p
367
  mask[1:], mask[0] = mask[:-1].clone(), False
368
  logits[sorted_i[mask]] = -float('Inf')
369
- text_token = torch.multinomial(F.softmax(logits, dim=-1), 1).item()
 
 
 
 
370
 
371
  if text_finished:
372
  text_token = args.get('enter_token_id', 201) if first_finished else args.get('pad_token_id', 0)
 
357
  past_kvs = out.past_key_values
358
 
359
  logits = out.logits[0, -1, :].clone().float() / (temperature + 1e-9)
360
+ logits = torch.nan_to_num(logits, nan=-100.0, posinf=-100.0, neginf=-100.0)
361
  if rp != 1.0:
362
  seen = list(set(input_ids[0].tolist()))
363
  score = logits[seen]
 
367
  mask = torch.cumsum(F.softmax(sorted_l, dim=-1), dim=-1) > top_p
368
  mask[1:], mask[0] = mask[:-1].clone(), False
369
  logits[sorted_i[mask]] = -float('Inf')
370
+ probs = F.softmax(logits, dim=-1)
371
+ probs = torch.nan_to_num(probs)
372
+ if probs.sum() <= 0:
373
+ probs = torch.ones_like(probs) / probs.shape[-1]
374
+ text_token = torch.multinomial(probs, 1).item()
375
 
376
  if text_finished:
377
  text_token = args.get('enter_token_id', 201) if first_finished else args.get('pad_token_id', 0)