File size: 7,613 Bytes
f0a4e91
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
# Modified from https://github.com/MeiGen-AI/InfiniteTalk/blob/main/src/audio_analysis/wav2vec2.py
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
# Modified from InfiniteTalk original implementation

import librosa
import torch
import torch.nn.functional as F
from diffusers.configuration_utils import ConfigMixin
from diffusers.loaders.single_file_model import FromOriginalModelMixin
from diffusers.models.modeling_utils import ModelMixin
from einops import rearrange
from transformers import Wav2Vec2Config, Wav2Vec2FeatureExtractor
from transformers import Wav2Vec2Model as TransformersWav2Vec2Model
from transformers.modeling_outputs import BaseModelOutput


def linear_interpolation(features, seq_len):
    """Linear interpolation for audio features."""
    features = features.transpose(1, 2)
    output_features = F.interpolate(features, size=seq_len, align_corners=True, mode='linear')
    return output_features.transpose(1, 2)


class Wav2Vec2Model(TransformersWav2Vec2Model):
    """

    Custom Wav2Vec2Model that supports seq_len parameter for time alignment.

    This matches the official InfiniteTalk implementation.

    """
    def __init__(self, config: Wav2Vec2Config):
        super().__init__(config)

    def forward(

        self,

        input_values,

        seq_len,

        attention_mask=None,

        mask_time_indices=None,

        output_attentions=None,

        output_hidden_states=None,

        return_dict=None,

    ):
        output_hidden_states = (
            output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
        )
        return_dict = return_dict if return_dict is not None else self.config.use_return_dict

        extract_features = self.feature_extractor(input_values)
        extract_features = extract_features.transpose(1, 2)
        extract_features = linear_interpolation(extract_features, seq_len=seq_len)

        if attention_mask is not None:
            # compute reduced attention_mask corresponding to feature vectors
            attention_mask = self._get_feature_vector_attention_mask(
                extract_features.shape[1], attention_mask, add_adapter=False
            )

        hidden_states, extract_features = self.feature_projection(extract_features)
        hidden_states = self._mask_hidden_states(
            hidden_states, mask_time_indices=mask_time_indices, attention_mask=attention_mask
        )

        encoder_outputs = self.encoder(
            hidden_states,
            attention_mask=attention_mask,
            output_attentions=output_attentions,
            output_hidden_states=output_hidden_states,
            return_dict=return_dict,
        )

        hidden_states = encoder_outputs[0]

        if self.adapter is not None:
            hidden_states = self.adapter(hidden_states)

        if not return_dict:
            return (hidden_states, ) + encoder_outputs[1:]
        return BaseModelOutput(
            last_hidden_state=hidden_states,
            hidden_states=encoder_outputs.hidden_states,
            attentions=encoder_outputs.attentions,
        )


class FlashHeadAudioEncoder(ModelMixin, ConfigMixin, FromOriginalModelMixin):
    """

    Audio encoder for InfiniteTalk model.

    

    Uses Wav2Vec2Model (not Wav2Vec2ForCTC) to extract audio features,

    matching the original InfiniteTalk implementation.

    """
    
    def __init__(self, pretrained_model_path="facebook/wav2vec2-base-960h", device='cpu'):
        super(FlashHeadAudioEncoder, self).__init__()
        
        # Load pretrained model
        self.feature_extractor = Wav2Vec2FeatureExtractor.from_pretrained(pretrained_model_path)
        self.model = Wav2Vec2Model.from_pretrained(pretrained_model_path, output_attentions=True)
        
        # Freeze feature extractor
        self.model.feature_extractor._freeze_parameters()
        
        self.model = self.model.to(device)
        self.model.eval()
        
        # Video frame rate
        self.video_rate = 25  # InfiniteTalk uses 25 fps
        
    def extract_audio_feat(

        self,

        audio_path,

        return_all_layers=True,

        sr=16000,

        video_length=None,

    ):
        """

        Extract audio features from audio file.

        

        Args:

            audio_path: Path to audio file

            return_all_layers: Whether to return all hidden states (default True for InfiniteTalk)

            sr: Sample rate (default 16000)

            video_length: Target video length in frames

            

        Returns:

            Audio features tensor

        """
        # Load audio
        audio_input, sample_rate = librosa.load(audio_path, sr=sr)
        
        # Calculate video_length if not provided
        if video_length is None:
            audio_duration = len(audio_input) / sr
            video_length = int(audio_duration * self.video_rate)
        
        # Extract features
        input_values = self.feature_extractor(
            audio_input, sampling_rate=sample_rate, return_tensors="pt"
        ).input_values
        
        # Inference
        with torch.no_grad():
            res = self.model(
                input_values.to(self.model.device),
                seq_len=video_length,  # Custom Wav2Vec2Model supports seq_len
                output_hidden_states=True
            )
            
            if return_all_layers:
                # Stack all hidden states (excluding embedding layer)
                feat = torch.stack(res.hidden_states[1:], dim=1).squeeze(0)
                feat = rearrange(feat, "b s d -> s b d")
            else:
                feat = res.hidden_states[-1]
        
        return feat

    def extract_audio_feat_without_file_load(

        self, 

        audio_array, 

        sample_rate, 

        return_all_layers=True, 

        video_length=None

    ):
        """

        Extract audio features from audio array (streaming mode).

        Matches official FlashHead preprocess_audio logic.

        

        Args:

            audio_array: Audio array (numpy or torch)

            sample_rate: Sample rate of the audio

            return_all_layers: Whether to return all hidden states

            video_length: Target video length in frames

            

        Returns:

            Audio features tensor [T, num_layers, dim]

        """
        # Convert to numpy if tensor
        if isinstance(audio_array, torch.Tensor):
            audio_array = audio_array.cpu().numpy()
        
        # Extract features
        input_values = self.feature_extractor(
            audio_array, sampling_rate=sample_rate, return_tensors="pt"
        ).input_values
        
        # Calculate video_length if not provided
        if video_length is None:
            audio_duration = len(audio_array) / sample_rate
            video_length = int(audio_duration * self.video_rate)
        
        # Inference
        with torch.no_grad():
            res = self.model(
                input_values.to(self.model.device),
                seq_len=video_length,
                output_hidden_states=True
            )
            
            if return_all_layers:
                feat = torch.stack(res.hidden_states[1:], dim=1).squeeze(0)
                feat = rearrange(feat, "b s d -> s b d")
            else:
                feat = res.hidden_states[-1]
        
        return feat