使用HuggingFace Transformer生成多文本样本时RuntimeError问题问询
Got it, let's break down why this error is happening and fix it so you can generate multiple samples smoothly.
What's Causing the Error?
When you set num_samples=5, your initial generated tensor has shape (5, context_length) (since you repeat the context 5 times). But in your loop, you're only grabbing logits for the first sample with outputs[0][0, -1, :], which produces a single token. When you unsqueeze this to (1,1) and try to concatenate it with generated (shape (5, seq_len)), PyTorch throws a dimension mismatch error—you can only concatenate along dimension 1 if all other dimensions have matching sizes (here, 5 vs 1 don't line up).
The Fix: Handle All Samples at Once
You need to process every sample in the batch simultaneously instead of just the first one. Here's the adjusted code with key fixes explained:
def sample_sequence( model, length, context, num_samples=1, temperature=1, top_k=0, top_p=0.9, repetition_penalty=1.0, device="cpu", ): context = torch.tensor(context, dtype=torch.long, device=device) context = context.unsqueeze(0).repeat(num_samples, 1) print('context_size', context.shape) generated = context print('context', context) with torch.no_grad(): for _ in trange(length): inputs = {"input_ids": generated} outputs = model(**inputs) # 1. Get logits for the last token of ALL samples in the batch next_token_logits = outputs[0][:, -1, :] / (temperature if temperature > 0 else 1.0) # 2. Apply repetition penalty PER sample (each sample has its own generated tokens) for sample_idx in range(num_samples): generated_tokens = set(generated[sample_idx].tolist()) for token in generated_tokens: next_token_logits[sample_idx, token] /= repetition_penalty filtered_logits = top_k_top_p_filtering(next_token_logits, top_k=top_k, top_p=top_p) if temperature == 0: # Greedy sampling: pick highest logit token for each sample next_token = torch.argmax(filtered_logits, dim=-1) else: # Multinomial sampling: pick one token per sample next_token = torch.multinomial(F.softmax(filtered_logits, dim=-1), num_samples=1).squeeze(1) # 3. Reshape to match generated's dimension for concatenation next_token = next_token.unsqueeze(1) # Now both tensors have matching first dimensions (num_samples) generated = torch.cat((generated, next_token), dim=1) return generated
Key Changes Explained
- Batch logits extraction: Using
outputs[0][:, -1, :]gives us a tensor of shape(num_samples, vocab_size), one set of logits for each sample in the batch. - Per-sample repetition penalty: We loop through each sample to adjust logits based on its own generated tokens, ensuring penalty is applied correctly to individual sequences.
- Batch token sampling: The sampling step now generates one token per sample, resulting in a
next_tokentensor of shape(num_samples,1)—this matches the first dimension ofgenerated, making concatenation along dim=1 valid.
内容的提问来源于stack exchange,提问作者Shamoon

