espnet2.aqa.ar_universa.ar_universa.ARUniversa
espnet2.aqa.ar_universa.ar_universa.ARUniversa
class espnet2.aqa.ar_universa.ar_universa.ARUniversa(input_size: int, metric2id: Dict[str, int], use_ref_audio: bool = True, use_ref_text: bool = True, embedding_size: int = 512, use_normalize: bool = True, audio_encoder_type: str = 'transformer', audio_encoder_params: Dict[str, float | int | bool | str] = {'attention_dropout_rate': 0.1, 'attention_heads': 4, 'concat_after': False, 'dropout_rate': 0.1, 'input_layer': 'linear', 'layer_drop_rate': 0.0, 'linear_units': 2048, 'normalize_before': True, 'num_blocks': 3, 'positional_dropout_rate': 0.1, 'positionwise_conv_kernel_size': 1, 'positionwise_layer_type': 'linear', 'qk_norm': False, 'use_flash_attn': False}, metric_vocab_size: int | None = None, metric_token_info: Dict[str, Any] | None = None, metric2type: Dict[str, str] | None = None, metric_pad_value: float = -100, metric_token_pad_value: int = 0, sequential_metrics: bool = True, vocab_size: int | None = None, ignore_id: int = -1, text_encoder_type: str = 'transformer', text_encoder_params: Dict[str, float | int | bool | str] = {'attention_dropout_rate': 0.1, 'attention_heads': 4, 'concat_after': False, 'dropout_rate': 0.1, 'input_layer': 'linear', 'layer_drop_rate': 0.0, 'linear_units': 2048, 'normalize_before': True, 'num_blocks': 3, 'positional_dropout_rate': 0.1, 'positionwise_conv_kernel_size': 1, 'positionwise_layer_type': 'linear', 'qk_norm': False, 'use_flash_attn': False}, cross_attention_type: str = 'multihead', cross_attention_params: Dict[str, float | int] = {'dropout_rate': 0.1, 'n_head': 4}, metric_decoder_params: Dict[str, float | int] = {'attention_heads': 4, 'concat_after': False, 'dropout_rate': 0.1, 'input_layer': 'embed', 'layer_drop_rate': 0.0, 'linear_units': 2048, 'normalize_before': True, 'num_blocks': 3, 'positional_dropout_rate': 0.1, 'qk_norm': False, 'self_attention_dropout_rate': 0.1, 'src_attention_dropout_rate': 0.1, 'use_flash_attn': False, 'use_output_layer': True}, use_rope_pos: bool = False, lsm_weight: float = 0.0, sym_sos: str = '<sos>', sym_eos: str = '<eos>', **kwargs)
Bases: AbsUniversa
Encode audio and optional references, then decode metric/value pairs.
Initialize ARECHO with the published architecture and token convention.
- Parameters:
- input_size (int) – Input feature size.
- metric2id (Dict *[*str , int ]) – Dictionary mapping metric names to IDs.
- use_ref_audio (bool) – Whether to use reference audio.
- use_ref_text (bool) – Whether to use reference text.
- embedding_size (int) – Embedding size for audio and text encoders.
- use_normalize (bool) – Whether to use normalization.
- audio_encoder_type (str) – Type of audio encoder.
- audio_encoder_params (Dict *[*str , Union *[*float , int , bool , str ] ]) – Parameters for audio encoder.
- metric_vocab_size (Optional *[*int ]) – Vocabulary size for metrics.
- metric_token_info (Optional *[*Dict *[*str , Any ] ]) – Information about metric tokens.
- metric2type (Optional *[*Dict *[*str , str ] ]) – Legacy config field. Metric types are defined by metric_token_info.
- metric_pad_value (float) – Legacy config field for regression metrics.
- metric_token_pad_value (int) – Padding value for metric tokens.
- sequential_metrics (bool) – Whether to use sequential metrics.
- vocab_size (Optional *[*int ]) – Vocabulary size for text encoder.
- ignore_id (int) – Ignore ID for padding in text encoder.
- text_encoder_type (str) – Type of text encoder.
- text_encoder_params (Dict *[*str , Union *[*float , int , bool , str ] ]) – Parameters for text encoder.
- cross_attention_type (str) – Type of cross attention module.
- cross_attention_params (Dict *[*str , Union *[*float , int ] ]) – Parameters for cross attention module.
- metric_decoder_params (Dict *[*str , Union *[*float , int ] ]) – Parameters for metric decoder module.
- use_rope_pos (bool) – Whether to use RoPE positional encoding.
- lsm_weight (float) – Label smoothing weight.
- sym_sos (str) – Legacy config field; ARECHO uses SOS ID 2.
- sym_eos (str) – Legacy config field; ARECHO uses EOS ID 3.
- **kwargs – Additional parameters.
encode(audio: Tensor, audio_lengths: Tensor, ref_audio: Tensor | None = None, ref_audio_lengths: Tensor | None = None, ref_text: Tensor | None = None, ref_text_lengths: Tensor | None = None, **kwargs) → Tuple[Tensor, Tensor]
Encode references without modifying caller-owned inputs.
Missing references contribute zero features in their configured slots.
forward(audio: Tensor, audio_lengths: Tensor, metrics: Dict[str, Tensor], ref_audio: Tensor | None = None, ref_audio_lengths: Tensor | None = None, ref_text: Tensor | None = None, ref_text_lengths: Tensor | None = None, **kwargs) → Tuple[Tensor, Dict[str, Tensor], Tensor]
Calculate outputs and return the loss tensor.
- Parameters:
- audio (torch.Tensor) – Input audio tensor (B, T).
- audio_lengths (torch.Tensor) – Length of audio tensor (B,).
- metrics (torch.Tensor) – Metrics tensor Dict[str, tensor (B,)].
- ref_audio (torch.Tensor) – Reference audio tensor (B, T).
- ref_audio_lengths (torch.Tensor) – Length of reference audio tensor (B,).
- ref_text (torch.Tensor) – Reference text tensor (B, U).
- ref_text_lengths (torch.Tensor) – Length of reference text tensor (B,).
- Returns: loss (torch.Tensor): Loss tensor. stats (Dict[str, torch.Tensor]): Statistics to be monitored. weight (torch.Tensor): Weight tensor.
- Return type: Tuple[torch.Tensor, Dict[str, torch.Tensor], torch.Tensor]
inference(audio: Tensor, audio_lengths: Tensor, ref_audio: Tensor | None = None, ref_audio_lengths: Tensor | None = None, ref_text: Tensor | None = None, ref_text_lengths: Tensor | None = None, **kwargs) → Dict[str, Any]
Return predicted output as a dict.
- Parameters:
- audio (torch.Tensor) – Input audio tensor (B, T).
- audio_lengths (torch.Tensor) – Length of audio tensor (B,).
- ref_audio (torch.Tensor) – Reference audio tensor (B, T).
- ref_audio_lengths (torch.Tensor) – Length of reference audio tensor (B,).
- ref_text (torch.Tensor) – Reference text tensor (B, U).
- ref_text_lengths (torch.Tensor) – Length of reference text tensor (B,).
- **kwargs – Additional parameters.
- Returns: Predicted output.
- Return type: Dict[str, torch.Tensor]
sequential_metrics = True
set_inference(beam_size: int, metric_list: List[str], skip_meta_label_score: bool, save_token_seq: bool = False, use_fixed_order: bool = False) → None
Set inference mode.
- Parameters:
- beam_size (int) – Beam size for beam search.
- metric_list (List *[*str ]) – List of metrics to predict.
- skip_meta_label_score (bool) – Whether to skip meta label score.
- save_token_seq (bool) – Whether to save token sequence.
- use_fixed_order (bool) – Decode metrics in the requested order.
