Team Ai
Modelpublic

HCKLab/BiBert-MultiTask-1

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

2�Vc;�@s�ddlZddlZddlmZddlmZmZmZddlmZm	Z	m3Z4mZddlmZddlm
Z
mZmZddlmZddlmZdd	lmZmZmZmZmZmZmZmZmZdd5lmZm Z m!Z!e!de�Gdd
�d
e��Z"dS)�N)�nn)�BCEWithLogitsLoss�CrossEntropyLoss�MSELoss)�List�Optional�Tuple�Union)�
BertTokenizer)�models�DataCollatorWithPadding�
AutoTokenizer)�SequenceClassifierOutput)�6BertConfig)	�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)NZ	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-1/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)r:r;r<r=r>r@rArBrr������)�loss�logits�
hidden_states�29attentions)r$�use_return_dictr%r(�torch�unique�tolistr,r-r�viewr#r+�append�stack�meanrrHrI)r.r9r:r;r<r=r>r?r@rArB�task_ids�outputsZ
pooled_outputZunique_task_ids_list�	loss_listrGZunique_task_idrFZtask_id_filterZloss_fct�outputr1r1r2�forward=sJ 3031$z%BertForSequenceClassification.forward)NNNNNNNNNNN)�__name__�32__module__�__qualname__rrr�formatrrrrrrrrrK�Tensor�boolr	rrV�
__classcell__r1r1)r0r2rs,33Lr)#rKr r�torch.nnrrr�typingrrrr	r34rrr
Ztransformers.modeling_outputsrZ+transformers.models.bert.configuration_bertrZ&transformers.models.bert.modeling_bertrrrrrrrrrZtransformers.file_utilsrrrrr1r1r1r2�<module>s,