CoolFace
Modelpublic

HCKLab/BiBert-MultiTask-2

sourceHugging Facemitupdated 4y agoView on Hugging Face
1likes7downloads
bert_for_sequence_classification.cpython-37.pyc35 linesDownload Raw Back to __pycache__
1B

2+�sc��@s�ddlZddlZddlmZddlmZddlmZmZmZddl	m3Z4ddlmZm
Z
mZmZmZmZmZmZmZddlmZmZmZede�Gd	d5�d6e��ZdS)�N)�nn)�CrossEntropyLoss)�Optional�Tuple�Union)�SequenceClassifierOutput)	�BertPreTrainedModel�BERT_INPUTS_DOCSTRING�_TOKENIZER_FOR_DOC�_CHECKPOINT_FOR_DOC�BERT_START_DOCSTRING�_CONFIG_FOR_DOC�_SEQ_CLASS_EXPECTED_OUTPUT�_SEQ_CLASS_EXPECTED_LOSS�	BertModel)�add_code_sample_docstrings�%add_start_docstrings_to_model_forward�add_start_docstringsz�7    Bert Model transformer with a sequence classification/regression head on top (a linear layer on top of the pooled8    output) e.g. for GLUE tasks.9    cs�eZdZ�fdd�Zee�d��eee	e10eee
d�d	eejeejeejeejeejeejeejeeeeeeeeeje11fd�dd���Z�ZS)12�BertForSequenceClassificationcs�t��t���|�di�|_||_t|�|_|j	dk	r>|j	n|j13}t�|�|_
t�|j|jdj�|_t�|j|jdj�|_|��dS)N�	tasks_mapr�)�super�__init__�transformers�PretrainedConfig�get�tasks�configr�bert�classifier_dropoutZhidden_dropout_probr�Dropout�dropout�Linear�hidden_size�14num_labels�classifier1�classifier2Zinit_weights)�selfr�kwargsr)�	__class__��?/content/BiBert-MultiTask-2/bert_for_sequence_classification.pyr!s15z&BertForSequenceClassification.__init__zbatch_size, sequence_length)�processor_class�16checkpoint�output_type�config_class�expected_output�
expected_lossN)�	input_ids�attention_mask�token_type_ids�position_ids�	head_mask�
inputs_embeds�labels�output_attentions�output_hidden_states�return_dict�returncCs2|17dk	r|18n|jj}19|j||||||||	|20d�	}|d}
|�|
�}
t�|���}g}d}x�|D]z}d}||k}|dkr�|�|
|�}n|dkr�|�|
|�}|dk	r^t	�}||�21d|j|j�||�22d��}|�
|�q^W|r�t�|���}|23�s|f|dd�}|dk	�r|f|S|St|||j|jd�S)a�24        labels (:obj:`torch.LongTensor` of shape :obj:`(batch_size,)`, `optional`):25            Labels for computing the sequence classification/regression loss. Indices should be in :obj:`[0, ...,26            config.num_labels - 1]`. If :obj:`config.num_labels == 1` a regression loss is computed (Mean-Square loss),27            If :obj:`config.num_labels > 1` a classification loss is computed (Cross-Entropy).28        N)r3r4r5r6r7r9r:r;rr������)�loss�logits�
hidden_states�29attentions)r�use_return_dictrr!�torch�unique�tolistr%r&r�viewrr$�append�stack�meanrrArB)r'r2r3r4r5r6r7r8r9r:r;�task_ids�outputsZ
pooled_outputZunique_task_ids_list�	loss_listr@Zunique_task_idr?Ztask_id_filterZloss_fct�outputr*r*r+�forward8sJ 3031$z%BertForSequenceClassification.forward)NNNNNNNNNNN)�__name__�32__module__�__qualname__rrr	�formatrr33rrr
rrrrD�Tensor�boolrrrO�
__classcell__r*r*)r)r+rs,34Lr)rDrr�torch.nnr�typingrrrZtransformers.modeling_outputsrZ&transformers.models.bert.modeling_bertrr	r35rrr
rrrZtransformers.file_utilsrrrrr*r*r*r+�<module>s,