File size: 15,944 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
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
# Modified from https://github.com/meituan-longcat/LongCat-Video/blob/main/longcat_video/audio_process/wav2vec2.py
import copy
import logging
import math
import os

import librosa
import numpy as np
import torch
import torch.nn as nn
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 Wav2Vec2Model_base
from transformers.activations import ACT2FN
from transformers.modeling_outputs import BaseModelOutput
from transformers.models.wav2vec2.modeling_wav2vec2 import (
    Wav2Vec2PositionalConvEmbedding, Wav2Vec2SamePadLayer)


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


def _Wav2Vec2PositionalConvEmbedding_init_hack_(self, config):
        super(Wav2Vec2PositionalConvEmbedding, self).__init__()
        self.conv = nn.Conv1d(
            config.hidden_size,
            config.hidden_size,
            kernel_size=config.num_conv_pos_embeddings,
            padding=config.num_conv_pos_embeddings // 2,
            groups=config.num_conv_pos_embedding_groups,
        )

        weight_norm = nn.utils.weight_norm
        if hasattr(nn.utils.parametrizations, "weight_norm"):
            weight_norm = nn.utils.parametrizations.weight_norm
        self.conv = weight_norm(self.conv, name="weight", dim=2)

        self.padding = Wav2Vec2SamePadLayer(config.num_conv_pos_embeddings)
        self.activation = ACT2FN[config.feat_extract_activation]


Wav2Vec2PositionalConvEmbedding.__init__ = _Wav2Vec2PositionalConvEmbedding_init_hack_


# the implementation of Wav2Vec2Model is borrowed from
# https://github.com/huggingface/transformers/blob/HEAD/src/transformers/models/wav2vec2/modeling_wav2vec2.py
# initialize our encoder with the pre-trained wav2vec 2.0 weights.
class Wav2Vec2Mode(Wav2Vec2Model_base):
    def __init__(self, config: Wav2Vec2Config):
        config.attn_implementation = "eager"
        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,

    ):
        self.config._attn_implementation = "eager"
        self.config.output_attentions = True

        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,
        )


    def feature_extract(

        self,

        input_values,

        seq_len,

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

        return extract_features

    def encode(

        self,

        extract_features,

        attention_mask=None,

        mask_time_indices=None,

        output_attentions=None,

        output_hidden_states=None,

        return_dict=None,

    ):
        self.config.output_attentions = True

        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

        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 Wav2Vec2ModelWrapper(nn.Module):
    def __init__(self, config_path, device='cuda', prefix='wav2vec2.'):
        super(Wav2Vec2ModelWrapper, self).__init__()

        config, model_kwargs = Wav2Vec2Config.from_pretrained(
                config_path,
                return_unused_kwargs=True,
                force_download=False,
                local_files_only=True,
            )

        model_path = os.path.join(config_path, 'pytorch_model.bin')
        state_dict = torch.load(model_path, map_location=device)

        config.name_or_path = config_path
        config = copy.deepcopy(config)  # We do not want to modify the config inplace in from_pretrained.
        # config = Wav2Vec2Mode._autoset_attn_implementation(config, use_flash_attention_2=False)

        # init model
        with torch.device('meta'):
            model = Wav2Vec2Mode(config)

        # load checkpoint
        logging.info(f'loading {model_path}')
        if prefix is not None:
            state_dict = {i.replace(prefix, ''):state_dict[i] for i in state_dict}
        
        model.tie_weights()
        m, u = model.load_state_dict(state_dict, assign=True, strict=False)
            
        model.tie_weights()
        model.eval()

        self.model = model
    
    @property
    def feature_extractor(self):
        return self.model.feature_extractor

    @property
    def dtype(self):
        return next(self.model.parameters()).dtype

    @property
    def device(self):
        return next(self.model.parameters()).device

    def forward(

        self,

        input_values,

        seq_len,

        attention_mask=None,

        mask_time_indices=None,

        output_attentions=None,

        output_hidden_states=None,

        return_dict=None,

    ):
        return self.model(
            input_values,
            seq_len,
            attention_mask=attention_mask,
            mask_time_indices=mask_time_indices,
            output_attentions=output_attentions,
            output_hidden_states=output_hidden_states,
            return_dict=return_dict,
        )
    
    def feature_extract(

        self,

        input_values,

        seq_len,

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

        return self.model.feature_extract(
            input_values,
            seq_len
        )

    def encode(

        self,

        extract_features,

        attention_mask=None,

        mask_time_indices=None,

        output_attentions=None,

        output_hidden_states=None,

        return_dict=None,

    ):

        return self.model.encode(
            extract_features,
            attention_mask=attention_mask,
            mask_time_indices=mask_time_indices,
            output_attentions=output_attentions,
            output_hidden_states=output_hidden_states,
            return_dict=return_dict,
        )


class LongCatVideoAudioEncoder(ModelMixin, ConfigMixin, FromOriginalModelMixin):
    """Audio encoder for LongCatVideo Avatar pipeline.

    

    This class provides a clean interface for audio feature extraction,

    similar to FantasyTalkingAudioEncoder but with LongCatVideo-specific

    audio preprocessing (loudness normalization, noise floor, transient smoothing).

    

    Uses existing Wav2Vec2ModelWrapper and Wav2Vec2FeatureExtractor internally.

    """
    
    def __init__(self, config_path, device='cpu', prefix='wav2vec2.'):
        super(LongCatVideoAudioEncoder, self).__init__()
        
        # Use existing Wav2Vec2ModelWrapper
        self.audio_encoder = Wav2Vec2ModelWrapper(config_path, device=device, prefix=prefix)
        
        # Use existing Wav2Vec2FeatureExtractor
        self.wav2vec_feature_extractor = Wav2Vec2FeatureExtractor.from_pretrained(config_path)
    
    @property
    def dtype(self):
        return self.audio_encoder.dtype

    @property
    def device(self):
        return self.audio_encoder.device
    
    def _loudness_norm(self, audio_array, sr=16000, lufs=-23, threshold=100):
        """Normalize audio loudness to target LUFS."""
        import pyloudnorm as pyln
        meter = pyln.Meter(sr)
        loudness = meter.integrated_loudness(audio_array)
        if abs(loudness) > threshold:
            return audio_array
        normalized_audio = pyln.normalize.loudness(audio_array, loudness, lufs)
        return normalized_audio

    def _add_noise_floor(self, audio, noise_db=-45):
        """Add noise floor to audio."""
        noise_amp = 10 ** (noise_db / 20)
        noise = np.random.randn(len(audio)) * noise_amp
        return audio + noise

    def _smooth_transients(self, audio, sr=16000):
        """Smooth audio transients using low-pass filter."""
        import scipy.signal as ss
        b, a = ss.butter(3, 3000 / (sr / 2))
        return ss.lfilter(b, a, audio)
    
    def _preprocess_audio(self, speech_array, sample_rate=16000):
        """Apply LongCatVideo-specific audio preprocessing."""
        speech_array = self._loudness_norm(speech_array, sample_rate)
        speech_array = self._add_noise_floor(speech_array)
        speech_array = self._smooth_transients(speech_array, sample_rate)
        return speech_array
    
    @torch.no_grad()
    def _extract_embedding(self, speech_array, sample_rate, num_frames, audio_stride=2):
        """Core method to extract audio embedding from preprocessed speech array.

        

        Args:

            speech_array: Preprocessed audio array.

            sample_rate: Audio sample rate.

            num_frames: Number of video frames.

            audio_stride: Audio stride for sliding window.

            

        Returns:

            Audio embeddings tensor of shape [1, num_frames, 5, 12, 768].

        """
        seq_len = int(audio_stride * num_frames)
        
        # wav2vec_feature_extractor
        audio_feature = np.squeeze(
            self.wav2vec_feature_extractor(speech_array, sampling_rate=sample_rate).input_values
        )
        audio_feature = torch.from_numpy(audio_feature).float().to(device=self.device, dtype=self.dtype)
        audio_feature = audio_feature.unsqueeze(0)

        # audio embedding using Wav2Vec2ModelWrapper
        embeddings = self.audio_encoder(audio_feature, seq_len=seq_len, output_hidden_states=True)

        audio_emb = torch.stack(embeddings.hidden_states[1:], dim=1).squeeze(0)
        audio_emb = rearrange(audio_emb, "b s d -> s b d").contiguous()  # T, 12, 768

        # Prepare audio embedding with sliding window
        indices = torch.arange(2 * 2 + 1) - 2  # [-2, -1, 0, 1, 2]
        audio_start_idx = 0
        audio_end_idx = audio_start_idx + audio_stride * num_frames
        
        center_indices = torch.arange(audio_start_idx, audio_end_idx, audio_stride).unsqueeze(1) + \
            indices.unsqueeze(0)
        center_indices = torch.clamp(center_indices, min=0, max=audio_emb.shape[0] - 1)
        audio_emb = audio_emb[center_indices][None, ...]  # [1, num_frames, 5, 12, 768]
        
        return audio_emb
    
    def extract_audio_feat(

        self, 

        audio_path, 

        num_frames=49, 

        fps=16, 

        sr=16000,

        audio_stride=2

    ):
        """Extract audio features from audio file.

        

        Args:

            audio_path: Path to audio file.

            num_frames: Number of video frames.

            fps: Video frames per second.

            sr: Audio sample rate.

            audio_stride: Audio stride for sliding window.

            

        Returns:

            Audio embeddings tensor of shape [1, num_frames, 5, 12, 768].

        """
        # Load audio
        speech_array, sample_rate = librosa.load(audio_path, sr=sr)
        
        # Pad audio to target length
        generate_duration = num_frames / fps
        source_duration = len(speech_array) / sample_rate
        added_sample_nums = math.ceil((generate_duration - source_duration) * sample_rate)
        if added_sample_nums > 0:
            speech_array = np.append(speech_array, [0.] * added_sample_nums)
        
        # Preprocess and extract embedding
        speech_array = self._preprocess_audio(speech_array, sample_rate)
        return self._extract_embedding(speech_array, sample_rate, num_frames, audio_stride)
    
    def extract_audio_feat_without_file_load(

        self, 

        audio_segment, 

        sample_rate, 

        num_frames=49, 

        audio_stride=2

    ):
        """Extract audio features from audio array without file loading.

        

        Args:

            audio_segment: Audio array (numpy array).

            sample_rate: Audio sample rate.

            num_frames: Number of video frames.

            audio_stride: Audio stride for sliding window.

            

        Returns:

            Audio embeddings tensor of shape [1, num_frames, 5, 12, 768].

        """
        # Preprocess and extract embedding
        speech_array = self._preprocess_audio(audio_segment, sample_rate)
        return self._extract_embedding(speech_array, sample_rate, num_frames, audio_stride)