You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

使用HuggingFace Transformer生成多文本样本时RuntimeError问题问询

How to Fix RuntimeError When Generating Multiple Text Samples with PyTorch

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_token tensor of shape (num_samples,1)—this matches the first dimension of generated, making concatenation along dim=1 valid.

内容的提问来源于stack exchange,提问作者Shamoon

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.08 07:43:03