ReVID / sample /jetengine_ext /layers /sampler.py
GuoruiSong's picture
Add files using upload-large-folder tool
b50f36e verified
Raw
History Blame Contribute Delete
1.59 kB
import torch
from torch.nn import functional as F
def top_k_logits(logits, k):
if k <= 0:
return logits
else:
values, _ = torch.topk(logits, k)
min_values = values[..., -1, None]
return torch.where(logits < min_values, torch.full_like(logits, float('-inf')), logits)
def top_p_logits(logits, p):
sorted_logits, sorted_indices = torch.sort(logits, descending=True)
cumulative_probs = torch.cumsum(F.softmax(sorted_logits, dim=-1), dim=-1)
sorted_mask = cumulative_probs > p
sorted_mask[..., 1:] = sorted_mask[..., :-1].clone()
sorted_mask[..., 0] = False
mask_indices = torch.scatter(torch.full_like(logits, False, dtype=torch.bool),
-1, sorted_indices, sorted_mask)
logits = logits.masked_fill(mask_indices, float('-inf'))
return logits
def sample_with_temperature_topk_topp(logits, temperature=1.0, top_k=0, top_p=1.0):
orig_shape = logits.shape[:-1] # [batch, block]
vocab_size = logits.shape[-1]
logits = logits.reshape(-1, vocab_size) # [batch*block, vocab]
if temperature != 1.0:
logits = logits / temperature
if top_k > 0:
logits = top_k_logits(logits, top_k)
if top_p < 1.0:
logits = top_p_logits(logits, top_p)
probs = F.softmax(logits, dim=-1) # shape: [batch*block, vocab]
assert probs.dim() == 2
token = torch.multinomial(probs, num_samples=1) # [batch*block, 1]
token_prob = torch.gather(probs, -1, token) # [batch*block, 1]
return token.view(*orig_shape), token_prob.view(*orig_shape)