Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 5 additions & 1 deletion src/Microsoft.ML.TorchSharp/NasBert/BertTaskType.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -4,6 +4,7 @@

using System;
using System.Collections.Generic;
using System.ComponentModel;
using System.Text;

namespace Microsoft.ML.TorchSharp.NasBert
Expand All@@ -17,7 +18,10 @@ public enum BertTaskType
MaskedLM = 1,
TextClassification = 2,
SentenceRegression = 3,
NameEntityRecognition = 4,
NamedEntityRecognition = 4,
[Obsolete("Please use NamedEntityRecognition instead", false)]
[EditorBrowsable(EditorBrowsableState.Never)]
NameEntityRecognition = NamedEntityRecognition,
QuestionAnswering = 5
}
}
12 changes: 6 additions & 6 deletions src/Microsoft.ML.TorchSharp/NasBert/NasBertTrainer.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -204,7 +204,7 @@ private protected override Module CreateModule(IChannel ch, IDataView input)
EnglishRoberta tokenizerModel = Tokenizer.RobertaModel();

NasBertModel model;
if (Parent.BertOptions.TaskType == BertTaskType.NameEntityRecognition)
if (Parent.BertOptions.TaskType == BertTaskType.NamedEntityRecognition)
model = new NerModel(Parent.BertOptions, tokenizerModel.PadIndex, tokenizerModel.SymbolsCount, Parent.Option.NumberOfClasses);
else
model = new ModelForPrediction(Parent.BertOptions, tokenizerModel.PadIndex, tokenizerModel.SymbolsCount, Parent.Option.NumberOfClasses);
Expand DownExpand Up@@ -268,7 +268,7 @@ private protected override torch.Tensor PrepareRowTensor()
private protected override void RunModelAndBackPropagate(ref List<Tensor> inputTensors, ref Tensor targetsTensor)
{
Tensor logits = default;
if (Parent.BertOptions.TaskType == BertTaskType.NameEntityRecognition)
if (Parent.BertOptions.TaskType == BertTaskType.NamedEntityRecognition)
{
int[,] lengthArray = new int[inputTensors.Count, 1];
for (int i = 0; i < inputTensors.Count; i++)
Expand All@@ -293,7 +293,7 @@ private protected override void RunModelAndBackPropagate(ref List<Tensor> inputT
torch.Tensor loss;
if (Parent.BertOptions.TaskType == BertTaskType.TextClassification)
loss = torch.nn.CrossEntropyLoss(reduction: Parent.BertOptions.Reduction).forward(logits, targetsTensor);
else if (Parent.BertOptions.TaskType == BertTaskType.NameEntityRecognition)
else if (Parent.BertOptions.TaskType == BertTaskType.NamedEntityRecognition)
{
targetsTensor = targetsTensor.@long().view(-1);
logits = logits.view(-1, logits.size(-1));
Expand DownExpand Up@@ -338,7 +338,7 @@ public override SchemaShape GetOutputSchema(SchemaShape inputSchema)
outColumns[Option.ScoreColumnName] = new SchemaShape.Column(Option.ScoreColumnName, SchemaShape.Column.VectorKind.Vector,
NumberDataViewType.Single, false, new SchemaShape(AnnotationUtils.AnnotationsForMulticlassScoreColumn(labelCol)));
}
else if (BertOptions.TaskType == BertTaskType.NameEntityRecognition)
else if (BertOptions.TaskType == BertTaskType.NamedEntityRecognition)
{
var metadata = new List<SchemaShape.Column>();
metadata.Add(new SchemaShape.Column(AnnotationUtils.Kinds.KeyValues, SchemaShape.Column.VectorKind.Vector,
Expand DownExpand Up@@ -387,7 +387,7 @@ private protected override void CheckInputSchema(SchemaShape inputSchema)
TextDataViewType.Instance.ToString(), sentenceCol2.GetTypeString());
}
}
else if (BertOptions.TaskType == BertTaskType.NameEntityRecognition)
else if (BertOptions.TaskType == BertTaskType.NamedEntityRecognition)
{
if (labelCol.ItemType != NumberDataViewType.UInt32)
throw Host.ExceptSchemaMismatch(nameof(inputSchema), "label", Option.LabelColumnName,
Expand DownExpand Up@@ -535,7 +535,7 @@ protected override DataViewSchema.DetachedColumn[] GetOutputColumnsCore()
info[1] = new DataViewSchema.DetachedColumn(Parent.Options.ScoreColumnName, new VectorDataViewType(NumberDataViewType.Single, Parent.Options.NumberOfClasses), meta.ToAnnotations());
return info;
}
else if (Parent.BertOptions.TaskType == BertTaskType.NameEntityRecognition)
else if (Parent.BertOptions.TaskType == BertTaskType.NamedEntityRecognition)
{
var info = new DataViewSchema.DetachedColumn[1];
var keyType = Parent.LabelColumn.Annotations.Schema.GetColumnOrNull(AnnotationUtils.Kinds.KeyValues)?.Type as VectorDataViewType;
Expand Down
8 changes: 4 additions & 4 deletions src/Microsoft.ML.TorchSharp/NasBert/NerTrainer.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -35,7 +35,7 @@ namespace Microsoft.ML.TorchSharp.NasBert
/// </summary>
/// <remarks>
/// <format type="text/markdown"><![CDATA[
/// To create this trainer, use [NER](xref:Microsoft.ML.TorchSharpCatalog.NameEntityRecognition(Microsoft.ML.MulticlassClassificationCatalog.MulticlassClassificationTrainers,System.String,System.String,System.String,Int32,Int32,Int32,Microsoft.ML.TorchSharp.NasBert.BertArchitecture,Microsoft.ML.IDataView)).
/// To create this trainer, use [NER](xref:Microsoft.ML.TorchSharpCatalog.NamedEntityRecognition(Microsoft.ML.MulticlassClassificationCatalog.MulticlassClassificationTrainers,System.String,System.String,System.String,Int32,Int32,Int32,Microsoft.ML.TorchSharp.NasBert.BertArchitecture,Microsoft.ML.IDataView)).
///
/// ### Input and Output Columns
/// The input label column data must be a Vector of [string](xref:Microsoft.ML.Data.TextDataViewType) type and the sentence columns must be of type<xref:Microsoft.ML.Data.TextDataViewType>.
Expand All@@ -54,7 +54,7 @@ namespace Microsoft.ML.TorchSharp.NasBert
/// | Exportable to ONNX | No |
///
/// ### Training Algorithm Details
/// Trains a Deep Neural Network(DNN) by leveraging an existing pre-trained NAS-BERT roBERTa model for the purpose of name entity recognition.
/// Trains a Deep Neural Network(DNN) by leveraging an existing pre-trained NAS-BERT roBERTa model for the purpose of named entity recognition.
/// ]]>
/// </format>
/// </remarks>
Expand DownExpand Up@@ -93,7 +93,7 @@ internal NerTrainer(IHostEnvironment env,
BatchSize = batchSize,
MaxEpoch = maxEpochs,
ValidationSet = validationSet,
TaskType = BertTaskType.NameEntityRecognition
TaskType = BertTaskType.NamedEntityRecognition
})
{
}
Expand DownExpand Up@@ -295,7 +295,7 @@ private static NerTransformer Create(IHostEnvironment env, ModelLoadContext ctx)

options.Sentence1ColumnName = ctx.LoadString();
options.Sentence2ColumnName = ctx.LoadStringOrNull();
options.TaskType = BertTaskType.NameEntityRecognition;
options.TaskType = BertTaskType.NamedEntityRecognition;

BinarySaver saver = new BinarySaver(env, new BinarySaver.Arguments());
DataViewType type;
Expand Down
47 changes: 43 additions & 4 deletions src/Microsoft.ML.TorchSharp/TorchSharpCatalog.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -4,6 +4,7 @@

using System;
using System.Collections.Generic;
using System.ComponentModel;
using System.Text;
using Microsoft.ML.Data;
using Microsoft.ML.TorchSharp.AutoFormerV2;
Expand DownExpand Up@@ -161,7 +162,45 @@ public static ObjectDetectionMetrics EvaluateObjectDetection(
}

/// <summary>
/// Fine tune a NAS-BERT model for Name Entity Recognition. The limit for any sentence is 512 tokens. Each word typically
/// Obsolete: please use the <see cref="NamedEntityRecognition(MulticlassClassificationCatalog.MulticlassClassificationTrainers, string, string, string, int, int, BertArchitecture, IDataView)"/> method instead
/// </summary>
/// <param name="catalog">The transform's catalog.</param>
/// <param name="labelColumnName">Name of the label column. Column should be a key type.</param>
/// <param name="outputColumnName">Name of the output column. It will be a key type. It is the predicted label.</param>
/// <param name="sentence1ColumnName">Name of the column for the first sentence.</param>
/// <param name="batchSize">Number of rows in the batch.</param>
/// <param name="maxEpochs">Maximum number of times to loop through your training set.</param>
/// <param name="architecture">Architecture for the model. Defaults to Roberta.</param>
/// <param name="validationSet">The validation set used while training to improve model quality.</param>
/// <returns></returns>
[Obsolete("Please use NamedEntityRecognition method instead", false)]
[EditorBrowsable(EditorBrowsableState.Never)]
public static NerTrainer NameEntityRecognition(
this MulticlassClassificationCatalog.MulticlassClassificationTrainers catalog,
string labelColumnName = DefaultColumnNames.Label,
string outputColumnName = DefaultColumnNames.PredictedLabel,
string sentence1ColumnName = "Sentence",
int batchSize = 32,
int maxEpochs = 10,
BertArchitecture architecture = BertArchitecture.Roberta,
IDataView validationSet = null)
=> NamedEntityRecognition(catalog, labelColumnName, outputColumnName, sentence1ColumnName, batchSize, maxEpochs, architecture, validationSet);

/// <summary>
/// Obsolete: please use the <see cref="NamedEntityRecognition(MulticlassClassificationCatalog.MulticlassClassificationTrainers, NerTrainer.NerOptions)"/> method instead
/// </summary>
/// <param name="catalog">The transform's catalog.</param>
/// <param name="options">The full set of advanced options.</param>
/// <returns></returns>
[Obsolete("Please use NamedEntityRecognition method instead", false)]
[EditorBrowsable(EditorBrowsableState.Never)]
public static NerTrainer NameEntityRecognition(
this MulticlassClassificationCatalog.MulticlassClassificationTrainers catalog,
NerTrainer.NerOptions options)
=> NamedEntityRecognition(catalog, options);

/// <summary>
/// Fine tune a NAS-BERT model for Named Entity Recognition. The limit for any sentence is 512 tokens. Each word typically
/// will map to a single token, and we automatically add 2 specical tokens (a start token and a separator token)
/// so in general this limit will be 510 words for all sentences.
/// </summary>
Expand All@@ -174,7 +213,7 @@ public static ObjectDetectionMetrics EvaluateObjectDetection(
/// <param name="architecture">Architecture for the model. Defaults to Roberta.</param>
/// <param name="validationSet">The validation set used while training to improve model quality.</param>
/// <returns></returns>
public static NerTrainer NameEntityRecognition(
public static NerTrainer NamedEntityRecognition(
this MulticlassClassificationCatalog.MulticlassClassificationTrainers catalog,
string labelColumnName = DefaultColumnNames.Label,
string outputColumnName = DefaultColumnNames.PredictedLabel,
Expand All@@ -186,12 +225,12 @@ public static NerTrainer NameEntityRecognition(
=> new NerTrainer(CatalogUtils.GetEnvironment(catalog), labelColumnName, outputColumnName, sentence1ColumnName, batchSize, maxEpochs, validationSet, architecture);

/// <summary>
/// Fine tune a Name Entity Recognition model.
/// Fine tune a Named Entity Recognition model.
/// </summary>
/// <param name="catalog">The transform's catalog.</param>
/// <param name="options">The full set of advanced options.</param>
/// <returns></returns>
public static NerTrainer NameEntityRecognition(
public static NerTrainer NamedEntityRecognition(
this MulticlassClassificationCatalog.MulticlassClassificationTrainers catalog,
NerTrainer.NerOptions options)
=> new NerTrainer(CatalogUtils.GetEnvironment(catalog), options);
Expand Down
2 changes: 1 addition & 1 deletion test/Microsoft.ML.Tests/NerTests.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -54,7 +54,7 @@ public void TestSimpleNer()
}));
var chain = new EstimatorChain<ITransformer>();
var estimator = chain.Append(ML.Transforms.Conversion.MapValueToKey("Label", keyData: labels))
.Append(ML.MulticlassClassification.Trainers.NameEntityRecognition(outputColumnName: "outputColumn"))
.Append(ML.MulticlassClassification.Trainers.NamedEntityRecognition(outputColumnName: "outputColumn"))
.Append(ML.Transforms.Conversion.MapKeyToValue("outputColumn"));

var estimatorSchema = estimator.GetOutputSchema(SchemaShape.Create(dataView.Schema));
Expand Down
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Add copy buttons to all
 blocks\n(function() {\n function addCopyButtons() {\n document.querySelectorAll('pre code').forEach(function(codeBlock) {\n if (codeBlock.parentElement.hasAttribute('data-copy-added')) return;\n codeBlock.parentElement.setAttribute('data-copy-added', 'true');\n \n var btn = document.createElement('button');\n btn.textContent = 'Copy';\n btn.style.cssText = 'position:absolute;top:4px;right:4px;padding:2px 8px;font-size:11px;background:#4ecdc4;border:none;border-radius:4px;color:#1a1a2e;cursor:pointer;opacity:0.7;transition:opacity 0.2s;';\n btn.onmouseover = function() { this.style.opacity = '1'; };\n btn.onmouseout = function() { this.style.opacity = '0.7'; };\n btn.onclick = function() {\n navigator.clipboard.writeText(codeBlock.textContent).then(function() {\n btn.textContent = 'Copied!';\n setTimeout(function() { btn.textContent = 'Copy'; }, 1500);\n });\n };\n codeBlock.parentElement.style.position = 'relative';\n codeBlock.parentElement.appendChild(btn);\n });\n }\n \n addCopyButtons();\n \n // Re-run on dynamic content\n var observer = new MutationObserver(addCopyButtons);\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "Add Copy Buttons to Code Blocks");
}
} catch(__e) { console.warn('[Userscript:Add Copy Buttons to Code Blocks]', __e); }
})();
(function(){
try {
var __m = "github.com";
var __re = new RegExp('^' + "github\\.com" + '
Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 5 additions & 1 deletion src/Microsoft.ML.TorchSharp/NasBert/BertTaskType.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -4,6 +4,7 @@

using System;
using System.Collections.Generic;
using System.ComponentModel;
using System.Text;

namespace Microsoft.ML.TorchSharp.NasBert
Expand All@@ -17,7 +18,10 @@ public enum BertTaskType
MaskedLM = 1,
TextClassification = 2,
SentenceRegression = 3,
NameEntityRecognition = 4,
NamedEntityRecognition = 4,
[Obsolete("Please use NamedEntityRecognition instead", false)]
[EditorBrowsable(EditorBrowsableState.Never)]
NameEntityRecognition = NamedEntityRecognition,
QuestionAnswering = 5
}
}
12 changes: 6 additions & 6 deletions src/Microsoft.ML.TorchSharp/NasBert/NasBertTrainer.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -204,7 +204,7 @@ private protected override Module CreateModule(IChannel ch, IDataView input)
EnglishRoberta tokenizerModel = Tokenizer.RobertaModel();

NasBertModel model;
if (Parent.BertOptions.TaskType == BertTaskType.NameEntityRecognition)
if (Parent.BertOptions.TaskType == BertTaskType.NamedEntityRecognition)
model = new NerModel(Parent.BertOptions, tokenizerModel.PadIndex, tokenizerModel.SymbolsCount, Parent.Option.NumberOfClasses);
else
model = new ModelForPrediction(Parent.BertOptions, tokenizerModel.PadIndex, tokenizerModel.SymbolsCount, Parent.Option.NumberOfClasses);
Expand DownExpand Up@@ -268,7 +268,7 @@ private protected override torch.Tensor PrepareRowTensor()
private protected override void RunModelAndBackPropagate(ref List<Tensor> inputTensors, ref Tensor targetsTensor)
{
Tensor logits = default;
if (Parent.BertOptions.TaskType == BertTaskType.NameEntityRecognition)
if (Parent.BertOptions.TaskType == BertTaskType.NamedEntityRecognition)
{
int[,] lengthArray = new int[inputTensors.Count, 1];
for (int i = 0; i < inputTensors.Count; i++)
Expand All@@ -293,7 +293,7 @@ private protected override void RunModelAndBackPropagate(ref List<Tensor> inputT
torch.Tensor loss;
if (Parent.BertOptions.TaskType == BertTaskType.TextClassification)
loss = torch.nn.CrossEntropyLoss(reduction: Parent.BertOptions.Reduction).forward(logits, targetsTensor);
else if (Parent.BertOptions.TaskType == BertTaskType.NameEntityRecognition)
else if (Parent.BertOptions.TaskType == BertTaskType.NamedEntityRecognition)
{
targetsTensor = targetsTensor.@long().view(-1);
logits = logits.view(-1, logits.size(-1));
Expand DownExpand Up@@ -338,7 +338,7 @@ public override SchemaShape GetOutputSchema(SchemaShape inputSchema)
outColumns[Option.ScoreColumnName] = new SchemaShape.Column(Option.ScoreColumnName, SchemaShape.Column.VectorKind.Vector,
NumberDataViewType.Single, false, new SchemaShape(AnnotationUtils.AnnotationsForMulticlassScoreColumn(labelCol)));
}
else if (BertOptions.TaskType == BertTaskType.NameEntityRecognition)
else if (BertOptions.TaskType == BertTaskType.NamedEntityRecognition)
{
var metadata = new List<SchemaShape.Column>();
metadata.Add(new SchemaShape.Column(AnnotationUtils.Kinds.KeyValues, SchemaShape.Column.VectorKind.Vector,
Expand DownExpand Up@@ -387,7 +387,7 @@ private protected override void CheckInputSchema(SchemaShape inputSchema)
TextDataViewType.Instance.ToString(), sentenceCol2.GetTypeString());
}
}
else if (BertOptions.TaskType == BertTaskType.NameEntityRecognition)
else if (BertOptions.TaskType == BertTaskType.NamedEntityRecognition)
{
if (labelCol.ItemType != NumberDataViewType.UInt32)
throw Host.ExceptSchemaMismatch(nameof(inputSchema), "label", Option.LabelColumnName,
Expand DownExpand Up@@ -535,7 +535,7 @@ protected override DataViewSchema.DetachedColumn[] GetOutputColumnsCore()
info[1] = new DataViewSchema.DetachedColumn(Parent.Options.ScoreColumnName, new VectorDataViewType(NumberDataViewType.Single, Parent.Options.NumberOfClasses), meta.ToAnnotations());
return info;
}
else if (Parent.BertOptions.TaskType == BertTaskType.NameEntityRecognition)
else if (Parent.BertOptions.TaskType == BertTaskType.NamedEntityRecognition)
{
var info = new DataViewSchema.DetachedColumn[1];
var keyType = Parent.LabelColumn.Annotations.Schema.GetColumnOrNull(AnnotationUtils.Kinds.KeyValues)?.Type as VectorDataViewType;
Expand Down
8 changes: 4 additions & 4 deletions src/Microsoft.ML.TorchSharp/NasBert/NerTrainer.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -35,7 +35,7 @@ namespace Microsoft.ML.TorchSharp.NasBert
/// </summary>
/// <remarks>
/// <format type="text/markdown"><![CDATA[
/// To create this trainer, use [NER](xref:Microsoft.ML.TorchSharpCatalog.NameEntityRecognition(Microsoft.ML.MulticlassClassificationCatalog.MulticlassClassificationTrainers,System.String,System.String,System.String,Int32,Int32,Int32,Microsoft.ML.TorchSharp.NasBert.BertArchitecture,Microsoft.ML.IDataView)).
/// To create this trainer, use [NER](xref:Microsoft.ML.TorchSharpCatalog.NamedEntityRecognition(Microsoft.ML.MulticlassClassificationCatalog.MulticlassClassificationTrainers,System.String,System.String,System.String,Int32,Int32,Int32,Microsoft.ML.TorchSharp.NasBert.BertArchitecture,Microsoft.ML.IDataView)).
///
/// ### Input and Output Columns
/// The input label column data must be a Vector of [string](xref:Microsoft.ML.Data.TextDataViewType) type and the sentence columns must be of type<xref:Microsoft.ML.Data.TextDataViewType>.
Expand All@@ -54,7 +54,7 @@ namespace Microsoft.ML.TorchSharp.NasBert
/// | Exportable to ONNX | No |
///
/// ### Training Algorithm Details
/// Trains a Deep Neural Network(DNN) by leveraging an existing pre-trained NAS-BERT roBERTa model for the purpose of name entity recognition.
/// Trains a Deep Neural Network(DNN) by leveraging an existing pre-trained NAS-BERT roBERTa model for the purpose of named entity recognition.
/// ]]>
/// </format>
/// </remarks>
Expand DownExpand Up@@ -93,7 +93,7 @@ internal NerTrainer(IHostEnvironment env,
BatchSize = batchSize,
MaxEpoch = maxEpochs,
ValidationSet = validationSet,
TaskType = BertTaskType.NameEntityRecognition
TaskType = BertTaskType.NamedEntityRecognition
})
{
}
Expand DownExpand Up@@ -295,7 +295,7 @@ private static NerTransformer Create(IHostEnvironment env, ModelLoadContext ctx)

options.Sentence1ColumnName = ctx.LoadString();
options.Sentence2ColumnName = ctx.LoadStringOrNull();
options.TaskType = BertTaskType.NameEntityRecognition;
options.TaskType = BertTaskType.NamedEntityRecognition;

BinarySaver saver = new BinarySaver(env, new BinarySaver.Arguments());
DataViewType type;
Expand Down
47 changes: 43 additions & 4 deletions src/Microsoft.ML.TorchSharp/TorchSharpCatalog.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -4,6 +4,7 @@

using System;
using System.Collections.Generic;
using System.ComponentModel;
using System.Text;
using Microsoft.ML.Data;
using Microsoft.ML.TorchSharp.AutoFormerV2;
Expand DownExpand Up@@ -161,7 +162,45 @@ public static ObjectDetectionMetrics EvaluateObjectDetection(
}

/// <summary>
/// Fine tune a NAS-BERT model for Name Entity Recognition. The limit for any sentence is 512 tokens. Each word typically
/// Obsolete: please use the <see cref="NamedEntityRecognition(MulticlassClassificationCatalog.MulticlassClassificationTrainers, string, string, string, int, int, BertArchitecture, IDataView)"/> method instead
/// </summary>
/// <param name="catalog">The transform's catalog.</param>
/// <param name="labelColumnName">Name of the label column. Column should be a key type.</param>
/// <param name="outputColumnName">Name of the output column. It will be a key type. It is the predicted label.</param>
/// <param name="sentence1ColumnName">Name of the column for the first sentence.</param>
/// <param name="batchSize">Number of rows in the batch.</param>
/// <param name="maxEpochs">Maximum number of times to loop through your training set.</param>
/// <param name="architecture">Architecture for the model. Defaults to Roberta.</param>
/// <param name="validationSet">The validation set used while training to improve model quality.</param>
/// <returns></returns>
[Obsolete("Please use NamedEntityRecognition method instead", false)]
[EditorBrowsable(EditorBrowsableState.Never)]
public static NerTrainer NameEntityRecognition(
this MulticlassClassificationCatalog.MulticlassClassificationTrainers catalog,
string labelColumnName = DefaultColumnNames.Label,
string outputColumnName = DefaultColumnNames.PredictedLabel,
string sentence1ColumnName = "Sentence",
int batchSize = 32,
int maxEpochs = 10,
BertArchitecture architecture = BertArchitecture.Roberta,
IDataView validationSet = null)
=> NamedEntityRecognition(catalog, labelColumnName, outputColumnName, sentence1ColumnName, batchSize, maxEpochs, architecture, validationSet);

/// <summary>
/// Obsolete: please use the <see cref="NamedEntityRecognition(MulticlassClassificationCatalog.MulticlassClassificationTrainers, NerTrainer.NerOptions)"/> method instead
/// </summary>
/// <param name="catalog">The transform's catalog.</param>
/// <param name="options">The full set of advanced options.</param>
/// <returns></returns>
[Obsolete("Please use NamedEntityRecognition method instead", false)]
[EditorBrowsable(EditorBrowsableState.Never)]
public static NerTrainer NameEntityRecognition(
this MulticlassClassificationCatalog.MulticlassClassificationTrainers catalog,
NerTrainer.NerOptions options)
=> NamedEntityRecognition(catalog, options);

/// <summary>
/// Fine tune a NAS-BERT model for Named Entity Recognition. The limit for any sentence is 512 tokens. Each word typically
/// will map to a single token, and we automatically add 2 specical tokens (a start token and a separator token)
/// so in general this limit will be 510 words for all sentences.
/// </summary>
Expand All@@ -174,7 +213,7 @@ public static ObjectDetectionMetrics EvaluateObjectDetection(
/// <param name="architecture">Architecture for the model. Defaults to Roberta.</param>
/// <param name="validationSet">The validation set used while training to improve model quality.</param>
/// <returns></returns>
public static NerTrainer NameEntityRecognition(
public static NerTrainer NamedEntityRecognition(
this MulticlassClassificationCatalog.MulticlassClassificationTrainers catalog,
string labelColumnName = DefaultColumnNames.Label,
string outputColumnName = DefaultColumnNames.PredictedLabel,
Expand All@@ -186,12 +225,12 @@ public static NerTrainer NameEntityRecognition(
=> new NerTrainer(CatalogUtils.GetEnvironment(catalog), labelColumnName, outputColumnName, sentence1ColumnName, batchSize, maxEpochs, validationSet, architecture);

/// <summary>
/// Fine tune a Name Entity Recognition model.
/// Fine tune a Named Entity Recognition model.
/// </summary>
/// <param name="catalog">The transform's catalog.</param>
/// <param name="options">The full set of advanced options.</param>
/// <returns></returns>
public static NerTrainer NameEntityRecognition(
public static NerTrainer NamedEntityRecognition(
this MulticlassClassificationCatalog.MulticlassClassificationTrainers catalog,
NerTrainer.NerOptions options)
=> new NerTrainer(CatalogUtils.GetEnvironment(catalog), options);
Expand Down
2 changes: 1 addition & 1 deletion test/Microsoft.ML.Tests/NerTests.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -54,7 +54,7 @@ public void TestSimpleNer()
}));
var chain = new EstimatorChain<ITransformer>();
var estimator = chain.Append(ML.Transforms.Conversion.MapValueToKey("Label", keyData: labels))
.Append(ML.MulticlassClassification.Trainers.NameEntityRecognition(outputColumnName: "outputColumn"))
.Append(ML.MulticlassClassification.Trainers.NamedEntityRecognition(outputColumnName: "outputColumn"))
.Append(ML.Transforms.Conversion.MapKeyToValue("outputColumn"));

var estimatorSchema = estimator.GetOutputSchema(SchemaShape.Create(dataView.Schema));
Expand Down
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Force GitHub README to respect dark mode\n(function() {\n var style = document.createElement('style');\n style.textContent = '\n .markdown-body {\n color-scheme: dark light;\n }\n .markdown-body pre { background: #161b22 !important; }\n .markdown-body code { background: rgba(110, 118, 129, 0.4) !important; }\n .markdown-body table th, .markdown-body table td { border-color: #30363d !important; }\n .markdown-body img { background: #0d1117; }\n .markdown-body blockquote { border-left-color: #8b949e; }\n .markdown-body hr { border-color: #30363d; }\n ';\n document.head.appendChild(style);\n})();", "GitHub Dark Mode README Fix"); } } catch(__e) { console.warn('[Userscript:GitHub Dark Mode README Fix]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 5 additions & 1 deletion src/Microsoft.ML.TorchSharp/NasBert/BertTaskType.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -4,6 +4,7 @@

using System;
using System.Collections.Generic;
using System.ComponentModel;
using System.Text;

namespace Microsoft.ML.TorchSharp.NasBert
Expand All@@ -17,7 +18,10 @@ public enum BertTaskType
MaskedLM = 1,
TextClassification = 2,
SentenceRegression = 3,
NameEntityRecognition = 4,
NamedEntityRecognition = 4,
[Obsolete("Please use NamedEntityRecognition instead", false)]
[EditorBrowsable(EditorBrowsableState.Never)]
NameEntityRecognition = NamedEntityRecognition,
QuestionAnswering = 5
}
}
12 changes: 6 additions & 6 deletions src/Microsoft.ML.TorchSharp/NasBert/NasBertTrainer.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -204,7 +204,7 @@ private protected override Module CreateModule(IChannel ch, IDataView input)
EnglishRoberta tokenizerModel = Tokenizer.RobertaModel();

NasBertModel model;
if (Parent.BertOptions.TaskType == BertTaskType.NameEntityRecognition)
if (Parent.BertOptions.TaskType == BertTaskType.NamedEntityRecognition)
model = new NerModel(Parent.BertOptions, tokenizerModel.PadIndex, tokenizerModel.SymbolsCount, Parent.Option.NumberOfClasses);
else
model = new ModelForPrediction(Parent.BertOptions, tokenizerModel.PadIndex, tokenizerModel.SymbolsCount, Parent.Option.NumberOfClasses);
Expand DownExpand Up@@ -268,7 +268,7 @@ private protected override torch.Tensor PrepareRowTensor()
private protected override void RunModelAndBackPropagate(ref List<Tensor> inputTensors, ref Tensor targetsTensor)
{
Tensor logits = default;
if (Parent.BertOptions.TaskType == BertTaskType.NameEntityRecognition)
if (Parent.BertOptions.TaskType == BertTaskType.NamedEntityRecognition)
{
int[,] lengthArray = new int[inputTensors.Count, 1];
for (int i = 0; i < inputTensors.Count; i++)
Expand All@@ -293,7 +293,7 @@ private protected override void RunModelAndBackPropagate(ref List<Tensor> inputT
torch.Tensor loss;
if (Parent.BertOptions.TaskType == BertTaskType.TextClassification)
loss = torch.nn.CrossEntropyLoss(reduction: Parent.BertOptions.Reduction).forward(logits, targetsTensor);
else if (Parent.BertOptions.TaskType == BertTaskType.NameEntityRecognition)
else if (Parent.BertOptions.TaskType == BertTaskType.NamedEntityRecognition)
{
targetsTensor = targetsTensor.@long().view(-1);
logits = logits.view(-1, logits.size(-1));
Expand DownExpand Up@@ -338,7 +338,7 @@ public override SchemaShape GetOutputSchema(SchemaShape inputSchema)
outColumns[Option.ScoreColumnName] = new SchemaShape.Column(Option.ScoreColumnName, SchemaShape.Column.VectorKind.Vector,
NumberDataViewType.Single, false, new SchemaShape(AnnotationUtils.AnnotationsForMulticlassScoreColumn(labelCol)));
}
else if (BertOptions.TaskType == BertTaskType.NameEntityRecognition)
else if (BertOptions.TaskType == BertTaskType.NamedEntityRecognition)
{
var metadata = new List<SchemaShape.Column>();
metadata.Add(new SchemaShape.Column(AnnotationUtils.Kinds.KeyValues, SchemaShape.Column.VectorKind.Vector,
Expand DownExpand Up@@ -387,7 +387,7 @@ private protected override void CheckInputSchema(SchemaShape inputSchema)
TextDataViewType.Instance.ToString(), sentenceCol2.GetTypeString());
}
}
else if (BertOptions.TaskType == BertTaskType.NameEntityRecognition)
else if (BertOptions.TaskType == BertTaskType.NamedEntityRecognition)
{
if (labelCol.ItemType != NumberDataViewType.UInt32)
throw Host.ExceptSchemaMismatch(nameof(inputSchema), "label", Option.LabelColumnName,
Expand DownExpand Up@@ -535,7 +535,7 @@ protected override DataViewSchema.DetachedColumn[] GetOutputColumnsCore()
info[1] = new DataViewSchema.DetachedColumn(Parent.Options.ScoreColumnName, new VectorDataViewType(NumberDataViewType.Single, Parent.Options.NumberOfClasses), meta.ToAnnotations());
return info;
}
else if (Parent.BertOptions.TaskType == BertTaskType.NameEntityRecognition)
else if (Parent.BertOptions.TaskType == BertTaskType.NamedEntityRecognition)
{
var info = new DataViewSchema.DetachedColumn[1];
var keyType = Parent.LabelColumn.Annotations.Schema.GetColumnOrNull(AnnotationUtils.Kinds.KeyValues)?.Type as VectorDataViewType;
Expand Down
8 changes: 4 additions & 4 deletions src/Microsoft.ML.TorchSharp/NasBert/NerTrainer.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -35,7 +35,7 @@ namespace Microsoft.ML.TorchSharp.NasBert
/// </summary>
/// <remarks>
/// <format type="text/markdown"><![CDATA[
/// To create this trainer, use [NER](xref:Microsoft.ML.TorchSharpCatalog.NameEntityRecognition(Microsoft.ML.MulticlassClassificationCatalog.MulticlassClassificationTrainers,System.String,System.String,System.String,Int32,Int32,Int32,Microsoft.ML.TorchSharp.NasBert.BertArchitecture,Microsoft.ML.IDataView)).
/// To create this trainer, use [NER](xref:Microsoft.ML.TorchSharpCatalog.NamedEntityRecognition(Microsoft.ML.MulticlassClassificationCatalog.MulticlassClassificationTrainers,System.String,System.String,System.String,Int32,Int32,Int32,Microsoft.ML.TorchSharp.NasBert.BertArchitecture,Microsoft.ML.IDataView)).
///
/// ### Input and Output Columns
/// The input label column data must be a Vector of [string](xref:Microsoft.ML.Data.TextDataViewType) type and the sentence columns must be of type<xref:Microsoft.ML.Data.TextDataViewType>.
Expand All@@ -54,7 +54,7 @@ namespace Microsoft.ML.TorchSharp.NasBert
/// | Exportable to ONNX | No |
///
/// ### Training Algorithm Details
/// Trains a Deep Neural Network(DNN) by leveraging an existing pre-trained NAS-BERT roBERTa model for the purpose of name entity recognition.
/// Trains a Deep Neural Network(DNN) by leveraging an existing pre-trained NAS-BERT roBERTa model for the purpose of named entity recognition.
/// ]]>
/// </format>
/// </remarks>
Expand DownExpand Up@@ -93,7 +93,7 @@ internal NerTrainer(IHostEnvironment env,
BatchSize = batchSize,
MaxEpoch = maxEpochs,
ValidationSet = validationSet,
TaskType = BertTaskType.NameEntityRecognition
TaskType = BertTaskType.NamedEntityRecognition
})
{
}
Expand DownExpand Up@@ -295,7 +295,7 @@ private static NerTransformer Create(IHostEnvironment env, ModelLoadContext ctx)

options.Sentence1ColumnName = ctx.LoadString();
options.Sentence2ColumnName = ctx.LoadStringOrNull();
options.TaskType = BertTaskType.NameEntityRecognition;
options.TaskType = BertTaskType.NamedEntityRecognition;

BinarySaver saver = new BinarySaver(env, new BinarySaver.Arguments());
DataViewType type;
Expand Down
47 changes: 43 additions & 4 deletions src/Microsoft.ML.TorchSharp/TorchSharpCatalog.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -4,6 +4,7 @@

using System;
using System.Collections.Generic;
using System.ComponentModel;
using System.Text;
using Microsoft.ML.Data;
using Microsoft.ML.TorchSharp.AutoFormerV2;
Expand DownExpand Up@@ -161,7 +162,45 @@ public static ObjectDetectionMetrics EvaluateObjectDetection(
}

/// <summary>
/// Fine tune a NAS-BERT model for Name Entity Recognition. The limit for any sentence is 512 tokens. Each word typically
/// Obsolete: please use the <see cref="NamedEntityRecognition(MulticlassClassificationCatalog.MulticlassClassificationTrainers, string, string, string, int, int, BertArchitecture, IDataView)"/> method instead
/// </summary>
/// <param name="catalog">The transform's catalog.</param>
/// <param name="labelColumnName">Name of the label column. Column should be a key type.</param>
/// <param name="outputColumnName">Name of the output column. It will be a key type. It is the predicted label.</param>
/// <param name="sentence1ColumnName">Name of the column for the first sentence.</param>
/// <param name="batchSize">Number of rows in the batch.</param>
/// <param name="maxEpochs">Maximum number of times to loop through your training set.</param>
/// <param name="architecture">Architecture for the model. Defaults to Roberta.</param>
/// <param name="validationSet">The validation set used while training to improve model quality.</param>
/// <returns></returns>
[Obsolete("Please use NamedEntityRecognition method instead", false)]
[EditorBrowsable(EditorBrowsableState.Never)]
public static NerTrainer NameEntityRecognition(
this MulticlassClassificationCatalog.MulticlassClassificationTrainers catalog,
string labelColumnName = DefaultColumnNames.Label,
string outputColumnName = DefaultColumnNames.PredictedLabel,
string sentence1ColumnName = "Sentence",
int batchSize = 32,
int maxEpochs = 10,
BertArchitecture architecture = BertArchitecture.Roberta,
IDataView validationSet = null)
=> NamedEntityRecognition(catalog, labelColumnName, outputColumnName, sentence1ColumnName, batchSize, maxEpochs, architecture, validationSet);

/// <summary>
/// Obsolete: please use the <see cref="NamedEntityRecognition(MulticlassClassificationCatalog.MulticlassClassificationTrainers, NerTrainer.NerOptions)"/> method instead
/// </summary>
/// <param name="catalog">The transform's catalog.</param>
/// <param name="options">The full set of advanced options.</param>
/// <returns></returns>
[Obsolete("Please use NamedEntityRecognition method instead", false)]
[EditorBrowsable(EditorBrowsableState.Never)]
public static NerTrainer NameEntityRecognition(
this MulticlassClassificationCatalog.MulticlassClassificationTrainers catalog,
NerTrainer.NerOptions options)
=> NamedEntityRecognition(catalog, options);

/// <summary>
/// Fine tune a NAS-BERT model for Named Entity Recognition. The limit for any sentence is 512 tokens. Each word typically
/// will map to a single token, and we automatically add 2 specical tokens (a start token and a separator token)
/// so in general this limit will be 510 words for all sentences.
/// </summary>
Expand All@@ -174,7 +213,7 @@ public static ObjectDetectionMetrics EvaluateObjectDetection(
/// <param name="architecture">Architecture for the model. Defaults to Roberta.</param>
/// <param name="validationSet">The validation set used while training to improve model quality.</param>
/// <returns></returns>
public static NerTrainer NameEntityRecognition(
public static NerTrainer NamedEntityRecognition(
this MulticlassClassificationCatalog.MulticlassClassificationTrainers catalog,
string labelColumnName = DefaultColumnNames.Label,
string outputColumnName = DefaultColumnNames.PredictedLabel,
Expand All@@ -186,12 +225,12 @@ public static NerTrainer NameEntityRecognition(
=> new NerTrainer(CatalogUtils.GetEnvironment(catalog), labelColumnName, outputColumnName, sentence1ColumnName, batchSize, maxEpochs, validationSet, architecture);

/// <summary>
/// Fine tune a Name Entity Recognition model.
/// Fine tune a Named Entity Recognition model.
/// </summary>
/// <param name="catalog">The transform's catalog.</param>
/// <param name="options">The full set of advanced options.</param>
/// <returns></returns>
public static NerTrainer NameEntityRecognition(
public static NerTrainer NamedEntityRecognition(
this MulticlassClassificationCatalog.MulticlassClassificationTrainers catalog,
NerTrainer.NerOptions options)
=> new NerTrainer(CatalogUtils.GetEnvironment(catalog), options);
Expand Down
2 changes: 1 addition & 1 deletion test/Microsoft.ML.Tests/NerTests.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -54,7 +54,7 @@ public void TestSimpleNer()
}));
var chain = new EstimatorChain<ITransformer>();
var estimator = chain.Append(ML.Transforms.Conversion.MapValueToKey("Label", keyData: labels))
.Append(ML.MulticlassClassification.Trainers.NameEntityRecognition(outputColumnName: "outputColumn"))
.Append(ML.MulticlassClassification.Trainers.NamedEntityRecognition(outputColumnName: "outputColumn"))
.Append(ML.Transforms.Conversion.MapKeyToValue("outputColumn"));

var estimatorSchema = estimator.GetOutputSchema(SchemaShape.Create(dataView.Schema));
Expand Down
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Highlight search terms from Google/DuckDuckGo/Bing referrer\n(function() {\n var ref = document.referrer;\n var terms = [];\n \n if (ref.includes('google.com') || ref.includes('duckduckgo.com') || ref.includes('bing.com')) {\n var url = new URL(ref);\n var q = url.searchParams.get('q') || url.searchParams.get('p');\n if (q) {\n terms = q.split(/\\s+/).filter(function(t) { return t.length > 2; });\n }\n }\n \n if (terms.length === 0) return;\n \n var style = document.createElement('style');\n style.textContent = '.userscript-highlight { background: #fbbf24; color: #1a1a2e; padding: 1px 3px; border-radius: 2px; }';\n document.head.appendChild(style);\n \n function highlight(node) {\n if (node.nodeType === 3) { // text node\n var text = node.textContent;\n var found = false;\n terms.forEach(function(term) {\n var regex = new RegExp('(' + term.replace(/[.*+?^${}()|[\\]\\\\]/g, '\\\\') + ')', 'gi');\n if (regex.test(text)) {\n found = true;\n var frag = document.createDocumentFragment();\n var parts = text.split(regex);\n parts.forEach(function(part, i) {\n if (i % 2 === 0) {\n frag.appendChild(document.createTextNode(part));\n } else {\n var span = document.createElement('span');\n span.className = 'userscript-highlight';\n span.textContent = part;\n frag.appendChild(span);\n }\n });\n node.parentNode.replaceChild(frag, node);\n }\n });\n } else if (node.nodeType === 1 && node.childNodes) { // element\n var skipTags = ['SCRIPT', 'STYLE', 'NOSCRIPT', 'TEXTAREA', 'INPUT', 'SELECT'];\n if (!skipTags.includes(node.tagName)) {\n Array.from(node.childNodes).forEach(highlight);\n }\n }\n }\n \n highlight(document.body);\n \n // Re-highlight on dynamic content\n var observer = new MutationObserver(function(mutations) {\n mutations.forEach(function(m) {\n m.addedNodes.forEach(function(node) {\n if (node.nodeType === 1 || node.nodeType === 3) highlight(node);\n });\n });\n });\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "Highlight Search Terms"); } } catch(__e) { console.warn('[Userscript:Highlight Search Terms]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 5 additions & 1 deletion src/Microsoft.ML.TorchSharp/NasBert/BertTaskType.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -4,6 +4,7 @@

using System;
using System.Collections.Generic;
using System.ComponentModel;
using System.Text;

namespace Microsoft.ML.TorchSharp.NasBert
Expand All@@ -17,7 +18,10 @@ public enum BertTaskType
MaskedLM = 1,
TextClassification = 2,
SentenceRegression = 3,
NameEntityRecognition = 4,
NamedEntityRecognition = 4,
[Obsolete("Please use NamedEntityRecognition instead", false)]
[EditorBrowsable(EditorBrowsableState.Never)]
NameEntityRecognition = NamedEntityRecognition,
QuestionAnswering = 5
}
}
12 changes: 6 additions & 6 deletions src/Microsoft.ML.TorchSharp/NasBert/NasBertTrainer.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -204,7 +204,7 @@ private protected override Module CreateModule(IChannel ch, IDataView input)
EnglishRoberta tokenizerModel = Tokenizer.RobertaModel();

NasBertModel model;
if (Parent.BertOptions.TaskType == BertTaskType.NameEntityRecognition)
if (Parent.BertOptions.TaskType == BertTaskType.NamedEntityRecognition)
model = new NerModel(Parent.BertOptions, tokenizerModel.PadIndex, tokenizerModel.SymbolsCount, Parent.Option.NumberOfClasses);
else
model = new ModelForPrediction(Parent.BertOptions, tokenizerModel.PadIndex, tokenizerModel.SymbolsCount, Parent.Option.NumberOfClasses);
Expand DownExpand Up@@ -268,7 +268,7 @@ private protected override torch.Tensor PrepareRowTensor()
private protected override void RunModelAndBackPropagate(ref List<Tensor> inputTensors, ref Tensor targetsTensor)
{
Tensor logits = default;
if (Parent.BertOptions.TaskType == BertTaskType.NameEntityRecognition)
if (Parent.BertOptions.TaskType == BertTaskType.NamedEntityRecognition)
{
int[,] lengthArray = new int[inputTensors.Count, 1];
for (int i = 0; i < inputTensors.Count; i++)
Expand All@@ -293,7 +293,7 @@ private protected override void RunModelAndBackPropagate(ref List<Tensor> inputT
torch.Tensor loss;
if (Parent.BertOptions.TaskType == BertTaskType.TextClassification)
loss = torch.nn.CrossEntropyLoss(reduction: Parent.BertOptions.Reduction).forward(logits, targetsTensor);
else if (Parent.BertOptions.TaskType == BertTaskType.NameEntityRecognition)
else if (Parent.BertOptions.TaskType == BertTaskType.NamedEntityRecognition)
{
targetsTensor = targetsTensor.@long().view(-1);
logits = logits.view(-1, logits.size(-1));
Expand DownExpand Up@@ -338,7 +338,7 @@ public override SchemaShape GetOutputSchema(SchemaShape inputSchema)
outColumns[Option.ScoreColumnName] = new SchemaShape.Column(Option.ScoreColumnName, SchemaShape.Column.VectorKind.Vector,
NumberDataViewType.Single, false, new SchemaShape(AnnotationUtils.AnnotationsForMulticlassScoreColumn(labelCol)));
}
else if (BertOptions.TaskType == BertTaskType.NameEntityRecognition)
else if (BertOptions.TaskType == BertTaskType.NamedEntityRecognition)
{
var metadata = new List<SchemaShape.Column>();
metadata.Add(new SchemaShape.Column(AnnotationUtils.Kinds.KeyValues, SchemaShape.Column.VectorKind.Vector,
Expand DownExpand Up@@ -387,7 +387,7 @@ private protected override void CheckInputSchema(SchemaShape inputSchema)
TextDataViewType.Instance.ToString(), sentenceCol2.GetTypeString());
}
}
else if (BertOptions.TaskType == BertTaskType.NameEntityRecognition)
else if (BertOptions.TaskType == BertTaskType.NamedEntityRecognition)
{
if (labelCol.ItemType != NumberDataViewType.UInt32)
throw Host.ExceptSchemaMismatch(nameof(inputSchema), "label", Option.LabelColumnName,
Expand DownExpand Up@@ -535,7 +535,7 @@ protected override DataViewSchema.DetachedColumn[] GetOutputColumnsCore()
info[1] = new DataViewSchema.DetachedColumn(Parent.Options.ScoreColumnName, new VectorDataViewType(NumberDataViewType.Single, Parent.Options.NumberOfClasses), meta.ToAnnotations());
return info;
}
else if (Parent.BertOptions.TaskType == BertTaskType.NameEntityRecognition)
else if (Parent.BertOptions.TaskType == BertTaskType.NamedEntityRecognition)
{
var info = new DataViewSchema.DetachedColumn[1];
var keyType = Parent.LabelColumn.Annotations.Schema.GetColumnOrNull(AnnotationUtils.Kinds.KeyValues)?.Type as VectorDataViewType;
Expand Down
8 changes: 4 additions & 4 deletions src/Microsoft.ML.TorchSharp/NasBert/NerTrainer.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -35,7 +35,7 @@ namespace Microsoft.ML.TorchSharp.NasBert
/// </summary>
/// <remarks>
/// <format type="text/markdown"><![CDATA[
/// To create this trainer, use [NER](xref:Microsoft.ML.TorchSharpCatalog.NameEntityRecognition(Microsoft.ML.MulticlassClassificationCatalog.MulticlassClassificationTrainers,System.String,System.String,System.String,Int32,Int32,Int32,Microsoft.ML.TorchSharp.NasBert.BertArchitecture,Microsoft.ML.IDataView)).
/// To create this trainer, use [NER](xref:Microsoft.ML.TorchSharpCatalog.NamedEntityRecognition(Microsoft.ML.MulticlassClassificationCatalog.MulticlassClassificationTrainers,System.String,System.String,System.String,Int32,Int32,Int32,Microsoft.ML.TorchSharp.NasBert.BertArchitecture,Microsoft.ML.IDataView)).
///
/// ### Input and Output Columns
/// The input label column data must be a Vector of [string](xref:Microsoft.ML.Data.TextDataViewType) type and the sentence columns must be of type<xref:Microsoft.ML.Data.TextDataViewType>.
Expand All@@ -54,7 +54,7 @@ namespace Microsoft.ML.TorchSharp.NasBert
/// | Exportable to ONNX | No |
///
/// ### Training Algorithm Details
/// Trains a Deep Neural Network(DNN) by leveraging an existing pre-trained NAS-BERT roBERTa model for the purpose of name entity recognition.
/// Trains a Deep Neural Network(DNN) by leveraging an existing pre-trained NAS-BERT roBERTa model for the purpose of named entity recognition.
/// ]]>
/// </format>
/// </remarks>
Expand DownExpand Up@@ -93,7 +93,7 @@ internal NerTrainer(IHostEnvironment env,
BatchSize = batchSize,
MaxEpoch = maxEpochs,
ValidationSet = validationSet,
TaskType = BertTaskType.NameEntityRecognition
TaskType = BertTaskType.NamedEntityRecognition
})
{
}
Expand DownExpand Up@@ -295,7 +295,7 @@ private static NerTransformer Create(IHostEnvironment env, ModelLoadContext ctx)

options.Sentence1ColumnName = ctx.LoadString();
options.Sentence2ColumnName = ctx.LoadStringOrNull();
options.TaskType = BertTaskType.NameEntityRecognition;
options.TaskType = BertTaskType.NamedEntityRecognition;

BinarySaver saver = new BinarySaver(env, new BinarySaver.Arguments());
DataViewType type;
Expand Down
47 changes: 43 additions & 4 deletions src/Microsoft.ML.TorchSharp/TorchSharpCatalog.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -4,6 +4,7 @@

using System;
using System.Collections.Generic;
using System.ComponentModel;
using System.Text;
using Microsoft.ML.Data;
using Microsoft.ML.TorchSharp.AutoFormerV2;
Expand DownExpand Up@@ -161,7 +162,45 @@ public static ObjectDetectionMetrics EvaluateObjectDetection(
}

/// <summary>
/// Fine tune a NAS-BERT model for Name Entity Recognition. The limit for any sentence is 512 tokens. Each word typically
/// Obsolete: please use the <see cref="NamedEntityRecognition(MulticlassClassificationCatalog.MulticlassClassificationTrainers, string, string, string, int, int, BertArchitecture, IDataView)"/> method instead
/// </summary>
/// <param name="catalog">The transform's catalog.</param>
/// <param name="labelColumnName">Name of the label column. Column should be a key type.</param>
/// <param name="outputColumnName">Name of the output column. It will be a key type. It is the predicted label.</param>
/// <param name="sentence1ColumnName">Name of the column for the first sentence.</param>
/// <param name="batchSize">Number of rows in the batch.</param>
/// <param name="maxEpochs">Maximum number of times to loop through your training set.</param>
/// <param name="architecture">Architecture for the model. Defaults to Roberta.</param>
/// <param name="validationSet">The validation set used while training to improve model quality.</param>
/// <returns></returns>
[Obsolete("Please use NamedEntityRecognition method instead", false)]
[EditorBrowsable(EditorBrowsableState.Never)]
public static NerTrainer NameEntityRecognition(
this MulticlassClassificationCatalog.MulticlassClassificationTrainers catalog,
string labelColumnName = DefaultColumnNames.Label,
string outputColumnName = DefaultColumnNames.PredictedLabel,
string sentence1ColumnName = "Sentence",
int batchSize = 32,
int maxEpochs = 10,
BertArchitecture architecture = BertArchitecture.Roberta,
IDataView validationSet = null)
=> NamedEntityRecognition(catalog, labelColumnName, outputColumnName, sentence1ColumnName, batchSize, maxEpochs, architecture, validationSet);

/// <summary>
/// Obsolete: please use the <see cref="NamedEntityRecognition(MulticlassClassificationCatalog.MulticlassClassificationTrainers, NerTrainer.NerOptions)"/> method instead
/// </summary>
/// <param name="catalog">The transform's catalog.</param>
/// <param name="options">The full set of advanced options.</param>
/// <returns></returns>
[Obsolete("Please use NamedEntityRecognition method instead", false)]
[EditorBrowsable(EditorBrowsableState.Never)]
public static NerTrainer NameEntityRecognition(
this MulticlassClassificationCatalog.MulticlassClassificationTrainers catalog,
NerTrainer.NerOptions options)
=> NamedEntityRecognition(catalog, options);

/// <summary>
/// Fine tune a NAS-BERT model for Named Entity Recognition. The limit for any sentence is 512 tokens. Each word typically
/// will map to a single token, and we automatically add 2 specical tokens (a start token and a separator token)
/// so in general this limit will be 510 words for all sentences.
/// </summary>
Expand All@@ -174,7 +213,7 @@ public static ObjectDetectionMetrics EvaluateObjectDetection(
/// <param name="architecture">Architecture for the model. Defaults to Roberta.</param>
/// <param name="validationSet">The validation set used while training to improve model quality.</param>
/// <returns></returns>
public static NerTrainer NameEntityRecognition(
public static NerTrainer NamedEntityRecognition(
this MulticlassClassificationCatalog.MulticlassClassificationTrainers catalog,
string labelColumnName = DefaultColumnNames.Label,
string outputColumnName = DefaultColumnNames.PredictedLabel,
Expand All@@ -186,12 +225,12 @@ public static NerTrainer NameEntityRecognition(
=> new NerTrainer(CatalogUtils.GetEnvironment(catalog), labelColumnName, outputColumnName, sentence1ColumnName, batchSize, maxEpochs, validationSet, architecture);

/// <summary>
/// Fine tune a Name Entity Recognition model.
/// Fine tune a Named Entity Recognition model.
/// </summary>
/// <param name="catalog">The transform's catalog.</param>
/// <param name="options">The full set of advanced options.</param>
/// <returns></returns>
public static NerTrainer NameEntityRecognition(
public static NerTrainer NamedEntityRecognition(
this MulticlassClassificationCatalog.MulticlassClassificationTrainers catalog,
NerTrainer.NerOptions options)
=> new NerTrainer(CatalogUtils.GetEnvironment(catalog), options);
Expand Down
2 changes: 1 addition & 1 deletion test/Microsoft.ML.Tests/NerTests.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -54,7 +54,7 @@ public void TestSimpleNer()
}));
var chain = new EstimatorChain<ITransformer>();
var estimator = chain.Append(ML.Transforms.Conversion.MapValueToKey("Label", keyData: labels))
.Append(ML.MulticlassClassification.Trainers.NameEntityRecognition(outputColumnName: "outputColumn"))
.Append(ML.MulticlassClassification.Trainers.NamedEntityRecognition(outputColumnName: "outputColumn"))
.Append(ML.Transforms.Conversion.MapKeyToValue("outputColumn"));

var estimatorSchema = estimator.GetOutputSchema(SchemaShape.Create(dataView.Schema));
Expand Down
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Strip utm_, fbclid, gclid, etc. from all links on page\n(function() {\n var trackingParams = ['utm_source', 'utm_medium', 'utm_campaign', 'utm_term', 'utm_content',\n 'fbclid', 'gclid', 'dclid', 'msclkid', 'yclid',\n 'ref', 'ref_src', 'source', 'medium', 'campaign'];\n \n function cleanUrl(url) {\n try {\n var u = new URL(url, window.location.origin);\n var changed = false;\n trackingParams.forEach(function(p) {\n if (u.searchParams.has(p)) {\n u.searchParams.delete(p);\n changed = true;\n }\n });\n return changed ? u.toString() : url;\n } catch (e) {\n return url;\n }\n }\n \n function cleanLinks() {\n document.querySelectorAll('a[href]').forEach(function(a) {\n var clean = cleanUrl(a.href);\n if (clean !== a.href) a.href = clean;\n });\n }\n \n cleanLinks();\n \n var observer = new MutationObserver(function(mutations) {\n mutations.forEach(function(m) {\n m.addedNodes.forEach(function(node) {\n if (node.nodeType === 1) {\n if (node.tagName === 'A') cleanLinks();\n node.querySelectorAll('a[href]').forEach(function(a) {\n var clean = cleanUrl(a.href);\n if (clean !== a.href) a.href = clean;\n });\n }\n });\n });\n });\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "Remove Tracking Parameters from Links"); } } catch(__e) { console.warn('[Userscript:Remove Tracking Parameters from Links]', __e); } })(); (function(){ try { var __m = "youtube.com"; var __re = new RegExp('^' + "youtube\\.com" + '
Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 5 additions & 1 deletion src/Microsoft.ML.TorchSharp/NasBert/BertTaskType.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -4,6 +4,7 @@

using System;
using System.Collections.Generic;
using System.ComponentModel;
using System.Text;

namespace Microsoft.ML.TorchSharp.NasBert
Expand All@@ -17,7 +18,10 @@ public enum BertTaskType
MaskedLM = 1,
TextClassification = 2,
SentenceRegression = 3,
NameEntityRecognition = 4,
NamedEntityRecognition = 4,
[Obsolete("Please use NamedEntityRecognition instead", false)]
[EditorBrowsable(EditorBrowsableState.Never)]
NameEntityRecognition = NamedEntityRecognition,
QuestionAnswering = 5
}
}
12 changes: 6 additions & 6 deletions src/Microsoft.ML.TorchSharp/NasBert/NasBertTrainer.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -204,7 +204,7 @@ private protected override Module CreateModule(IChannel ch, IDataView input)
EnglishRoberta tokenizerModel = Tokenizer.RobertaModel();

NasBertModel model;
if (Parent.BertOptions.TaskType == BertTaskType.NameEntityRecognition)
if (Parent.BertOptions.TaskType == BertTaskType.NamedEntityRecognition)
model = new NerModel(Parent.BertOptions, tokenizerModel.PadIndex, tokenizerModel.SymbolsCount, Parent.Option.NumberOfClasses);
else
model = new ModelForPrediction(Parent.BertOptions, tokenizerModel.PadIndex, tokenizerModel.SymbolsCount, Parent.Option.NumberOfClasses);
Expand DownExpand Up@@ -268,7 +268,7 @@ private protected override torch.Tensor PrepareRowTensor()
private protected override void RunModelAndBackPropagate(ref List<Tensor> inputTensors, ref Tensor targetsTensor)
{
Tensor logits = default;
if (Parent.BertOptions.TaskType == BertTaskType.NameEntityRecognition)
if (Parent.BertOptions.TaskType == BertTaskType.NamedEntityRecognition)
{
int[,] lengthArray = new int[inputTensors.Count, 1];
for (int i = 0; i < inputTensors.Count; i++)
Expand All@@ -293,7 +293,7 @@ private protected override void RunModelAndBackPropagate(ref List<Tensor> inputT
torch.Tensor loss;
if (Parent.BertOptions.TaskType == BertTaskType.TextClassification)
loss = torch.nn.CrossEntropyLoss(reduction: Parent.BertOptions.Reduction).forward(logits, targetsTensor);
else if (Parent.BertOptions.TaskType == BertTaskType.NameEntityRecognition)
else if (Parent.BertOptions.TaskType == BertTaskType.NamedEntityRecognition)
{
targetsTensor = targetsTensor.@long().view(-1);
logits = logits.view(-1, logits.size(-1));
Expand DownExpand Up@@ -338,7 +338,7 @@ public override SchemaShape GetOutputSchema(SchemaShape inputSchema)
outColumns[Option.ScoreColumnName] = new SchemaShape.Column(Option.ScoreColumnName, SchemaShape.Column.VectorKind.Vector,
NumberDataViewType.Single, false, new SchemaShape(AnnotationUtils.AnnotationsForMulticlassScoreColumn(labelCol)));
}
else if (BertOptions.TaskType == BertTaskType.NameEntityRecognition)
else if (BertOptions.TaskType == BertTaskType.NamedEntityRecognition)
{
var metadata = new List<SchemaShape.Column>();
metadata.Add(new SchemaShape.Column(AnnotationUtils.Kinds.KeyValues, SchemaShape.Column.VectorKind.Vector,
Expand DownExpand Up@@ -387,7 +387,7 @@ private protected override void CheckInputSchema(SchemaShape inputSchema)
TextDataViewType.Instance.ToString(), sentenceCol2.GetTypeString());
}
}
else if (BertOptions.TaskType == BertTaskType.NameEntityRecognition)
else if (BertOptions.TaskType == BertTaskType.NamedEntityRecognition)
{
if (labelCol.ItemType != NumberDataViewType.UInt32)
throw Host.ExceptSchemaMismatch(nameof(inputSchema), "label", Option.LabelColumnName,
Expand DownExpand Up@@ -535,7 +535,7 @@ protected override DataViewSchema.DetachedColumn[] GetOutputColumnsCore()
info[1] = new DataViewSchema.DetachedColumn(Parent.Options.ScoreColumnName, new VectorDataViewType(NumberDataViewType.Single, Parent.Options.NumberOfClasses), meta.ToAnnotations());
return info;
}
else if (Parent.BertOptions.TaskType == BertTaskType.NameEntityRecognition)
else if (Parent.BertOptions.TaskType == BertTaskType.NamedEntityRecognition)
{
var info = new DataViewSchema.DetachedColumn[1];
var keyType = Parent.LabelColumn.Annotations.Schema.GetColumnOrNull(AnnotationUtils.Kinds.KeyValues)?.Type as VectorDataViewType;
Expand Down
8 changes: 4 additions & 4 deletions src/Microsoft.ML.TorchSharp/NasBert/NerTrainer.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -35,7 +35,7 @@ namespace Microsoft.ML.TorchSharp.NasBert
/// </summary>
/// <remarks>
/// <format type="text/markdown"><![CDATA[
/// To create this trainer, use [NER](xref:Microsoft.ML.TorchSharpCatalog.NameEntityRecognition(Microsoft.ML.MulticlassClassificationCatalog.MulticlassClassificationTrainers,System.String,System.String,System.String,Int32,Int32,Int32,Microsoft.ML.TorchSharp.NasBert.BertArchitecture,Microsoft.ML.IDataView)).
/// To create this trainer, use [NER](xref:Microsoft.ML.TorchSharpCatalog.NamedEntityRecognition(Microsoft.ML.MulticlassClassificationCatalog.MulticlassClassificationTrainers,System.String,System.String,System.String,Int32,Int32,Int32,Microsoft.ML.TorchSharp.NasBert.BertArchitecture,Microsoft.ML.IDataView)).
///
/// ### Input and Output Columns
/// The input label column data must be a Vector of [string](xref:Microsoft.ML.Data.TextDataViewType) type and the sentence columns must be of type<xref:Microsoft.ML.Data.TextDataViewType>.
Expand All@@ -54,7 +54,7 @@ namespace Microsoft.ML.TorchSharp.NasBert
/// | Exportable to ONNX | No |
///
/// ### Training Algorithm Details
/// Trains a Deep Neural Network(DNN) by leveraging an existing pre-trained NAS-BERT roBERTa model for the purpose of name entity recognition.
/// Trains a Deep Neural Network(DNN) by leveraging an existing pre-trained NAS-BERT roBERTa model for the purpose of named entity recognition.
/// ]]>
/// </format>
/// </remarks>
Expand DownExpand Up@@ -93,7 +93,7 @@ internal NerTrainer(IHostEnvironment env,
BatchSize = batchSize,
MaxEpoch = maxEpochs,
ValidationSet = validationSet,
TaskType = BertTaskType.NameEntityRecognition
TaskType = BertTaskType.NamedEntityRecognition
})
{
}
Expand DownExpand Up@@ -295,7 +295,7 @@ private static NerTransformer Create(IHostEnvironment env, ModelLoadContext ctx)

options.Sentence1ColumnName = ctx.LoadString();
options.Sentence2ColumnName = ctx.LoadStringOrNull();
options.TaskType = BertTaskType.NameEntityRecognition;
options.TaskType = BertTaskType.NamedEntityRecognition;

BinarySaver saver = new BinarySaver(env, new BinarySaver.Arguments());
DataViewType type;
Expand Down
47 changes: 43 additions & 4 deletions src/Microsoft.ML.TorchSharp/TorchSharpCatalog.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -4,6 +4,7 @@

using System;
using System.Collections.Generic;
using System.ComponentModel;
using System.Text;
using Microsoft.ML.Data;
using Microsoft.ML.TorchSharp.AutoFormerV2;
Expand DownExpand Up@@ -161,7 +162,45 @@ public static ObjectDetectionMetrics EvaluateObjectDetection(
}

/// <summary>
/// Fine tune a NAS-BERT model for Name Entity Recognition. The limit for any sentence is 512 tokens. Each word typically
/// Obsolete: please use the <see cref="NamedEntityRecognition(MulticlassClassificationCatalog.MulticlassClassificationTrainers, string, string, string, int, int, BertArchitecture, IDataView)"/> method instead
/// </summary>
/// <param name="catalog">The transform's catalog.</param>
/// <param name="labelColumnName">Name of the label column. Column should be a key type.</param>
/// <param name="outputColumnName">Name of the output column. It will be a key type. It is the predicted label.</param>
/// <param name="sentence1ColumnName">Name of the column for the first sentence.</param>
/// <param name="batchSize">Number of rows in the batch.</param>
/// <param name="maxEpochs">Maximum number of times to loop through your training set.</param>
/// <param name="architecture">Architecture for the model. Defaults to Roberta.</param>
/// <param name="validationSet">The validation set used while training to improve model quality.</param>
/// <returns></returns>
[Obsolete("Please use NamedEntityRecognition method instead", false)]
[EditorBrowsable(EditorBrowsableState.Never)]
public static NerTrainer NameEntityRecognition(
this MulticlassClassificationCatalog.MulticlassClassificationTrainers catalog,
string labelColumnName = DefaultColumnNames.Label,
string outputColumnName = DefaultColumnNames.PredictedLabel,
string sentence1ColumnName = "Sentence",
int batchSize = 32,
int maxEpochs = 10,
BertArchitecture architecture = BertArchitecture.Roberta,
IDataView validationSet = null)
=> NamedEntityRecognition(catalog, labelColumnName, outputColumnName, sentence1ColumnName, batchSize, maxEpochs, architecture, validationSet);

/// <summary>
/// Obsolete: please use the <see cref="NamedEntityRecognition(MulticlassClassificationCatalog.MulticlassClassificationTrainers, NerTrainer.NerOptions)"/> method instead
/// </summary>
/// <param name="catalog">The transform's catalog.</param>
/// <param name="options">The full set of advanced options.</param>
/// <returns></returns>
[Obsolete("Please use NamedEntityRecognition method instead", false)]
[EditorBrowsable(EditorBrowsableState.Never)]
public static NerTrainer NameEntityRecognition(
this MulticlassClassificationCatalog.MulticlassClassificationTrainers catalog,
NerTrainer.NerOptions options)
=> NamedEntityRecognition(catalog, options);

/// <summary>
/// Fine tune a NAS-BERT model for Named Entity Recognition. The limit for any sentence is 512 tokens. Each word typically
/// will map to a single token, and we automatically add 2 specical tokens (a start token and a separator token)
/// so in general this limit will be 510 words for all sentences.
/// </summary>
Expand All@@ -174,7 +213,7 @@ public static ObjectDetectionMetrics EvaluateObjectDetection(
/// <param name="architecture">Architecture for the model. Defaults to Roberta.</param>
/// <param name="validationSet">The validation set used while training to improve model quality.</param>
/// <returns></returns>
public static NerTrainer NameEntityRecognition(
public static NerTrainer NamedEntityRecognition(
this MulticlassClassificationCatalog.MulticlassClassificationTrainers catalog,
string labelColumnName = DefaultColumnNames.Label,
string outputColumnName = DefaultColumnNames.PredictedLabel,
Expand All@@ -186,12 +225,12 @@ public static NerTrainer NameEntityRecognition(
=> new NerTrainer(CatalogUtils.GetEnvironment(catalog), labelColumnName, outputColumnName, sentence1ColumnName, batchSize, maxEpochs, validationSet, architecture);

/// <summary>
/// Fine tune a Name Entity Recognition model.
/// Fine tune a Named Entity Recognition model.
/// </summary>
/// <param name="catalog">The transform's catalog.</param>
/// <param name="options">The full set of advanced options.</param>
/// <returns></returns>
public static NerTrainer NameEntityRecognition(
public static NerTrainer NamedEntityRecognition(
this MulticlassClassificationCatalog.MulticlassClassificationTrainers catalog,
NerTrainer.NerOptions options)
=> new NerTrainer(CatalogUtils.GetEnvironment(catalog), options);
Expand Down
2 changes: 1 addition & 1 deletion test/Microsoft.ML.Tests/NerTests.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -54,7 +54,7 @@ public void TestSimpleNer()
}));
var chain = new EstimatorChain<ITransformer>();
var estimator = chain.Append(ML.Transforms.Conversion.MapValueToKey("Label", keyData: labels))
.Append(ML.MulticlassClassification.Trainers.NameEntityRecognition(outputColumnName: "outputColumn"))
.Append(ML.MulticlassClassification.Trainers.NamedEntityRecognition(outputColumnName: "outputColumn"))
.Append(ML.Transforms.Conversion.MapKeyToValue("outputColumn"));

var estimatorSchema = estimator.GetOutputSchema(SchemaShape.Create(dataView.Schema));
Expand Down
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Auto-enable theater mode on YouTube\n(function() {\n function tryTheater() {\n var btn = document.querySelector('button[aria-label=\"Theater mode\"], ytd-player #player button[title=\"Theater mode\"]');\n if (btn && !btn.classList.contains('activated')) {\n btn.click();\n }\n }\n \n // Try immediately\n tryTheater();\n \n // Try after navigation (SPA)\n var lastUrl = location.href;\n setInterval(function() {\n if (location.href !== lastUrl) {\n lastUrl = location.href;\n setTimeout(tryTheater, 500);\n }\n }, 1000);\n \n // Also try on player load\n var observer = new MutationObserver(tryTheater);\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "YouTube Theater Mode Default"); } } catch(__e) { console.warn('[Userscript:YouTube Theater Mode Default]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 5 additions & 1 deletion src/Microsoft.ML.TorchSharp/NasBert/BertTaskType.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -4,6 +4,7 @@

using System;
using System.Collections.Generic;
using System.ComponentModel;
using System.Text;

namespace Microsoft.ML.TorchSharp.NasBert
Expand All@@ -17,7 +18,10 @@ public enum BertTaskType
MaskedLM = 1,
TextClassification = 2,
SentenceRegression = 3,
NameEntityRecognition = 4,
NamedEntityRecognition = 4,
[Obsolete("Please use NamedEntityRecognition instead", false)]
[EditorBrowsable(EditorBrowsableState.Never)]
NameEntityRecognition = NamedEntityRecognition,
QuestionAnswering = 5
}
}
12 changes: 6 additions & 6 deletions src/Microsoft.ML.TorchSharp/NasBert/NasBertTrainer.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -204,7 +204,7 @@ private protected override Module CreateModule(IChannel ch, IDataView input)
EnglishRoberta tokenizerModel = Tokenizer.RobertaModel();

NasBertModel model;
if (Parent.BertOptions.TaskType == BertTaskType.NameEntityRecognition)
if (Parent.BertOptions.TaskType == BertTaskType.NamedEntityRecognition)
model = new NerModel(Parent.BertOptions, tokenizerModel.PadIndex, tokenizerModel.SymbolsCount, Parent.Option.NumberOfClasses);
else
model = new ModelForPrediction(Parent.BertOptions, tokenizerModel.PadIndex, tokenizerModel.SymbolsCount, Parent.Option.NumberOfClasses);
Expand DownExpand Up@@ -268,7 +268,7 @@ private protected override torch.Tensor PrepareRowTensor()
private protected override void RunModelAndBackPropagate(ref List<Tensor> inputTensors, ref Tensor targetsTensor)
{
Tensor logits = default;
if (Parent.BertOptions.TaskType == BertTaskType.NameEntityRecognition)
if (Parent.BertOptions.TaskType == BertTaskType.NamedEntityRecognition)
{
int[,] lengthArray = new int[inputTensors.Count, 1];
for (int i = 0; i < inputTensors.Count; i++)
Expand All@@ -293,7 +293,7 @@ private protected override void RunModelAndBackPropagate(ref List<Tensor> inputT
torch.Tensor loss;
if (Parent.BertOptions.TaskType == BertTaskType.TextClassification)
loss = torch.nn.CrossEntropyLoss(reduction: Parent.BertOptions.Reduction).forward(logits, targetsTensor);
else if (Parent.BertOptions.TaskType == BertTaskType.NameEntityRecognition)
else if (Parent.BertOptions.TaskType == BertTaskType.NamedEntityRecognition)
{
targetsTensor = targetsTensor.@long().view(-1);
logits = logits.view(-1, logits.size(-1));
Expand DownExpand Up@@ -338,7 +338,7 @@ public override SchemaShape GetOutputSchema(SchemaShape inputSchema)
outColumns[Option.ScoreColumnName] = new SchemaShape.Column(Option.ScoreColumnName, SchemaShape.Column.VectorKind.Vector,
NumberDataViewType.Single, false, new SchemaShape(AnnotationUtils.AnnotationsForMulticlassScoreColumn(labelCol)));
}
else if (BertOptions.TaskType == BertTaskType.NameEntityRecognition)
else if (BertOptions.TaskType == BertTaskType.NamedEntityRecognition)
{
var metadata = new List<SchemaShape.Column>();
metadata.Add(new SchemaShape.Column(AnnotationUtils.Kinds.KeyValues, SchemaShape.Column.VectorKind.Vector,
Expand DownExpand Up@@ -387,7 +387,7 @@ private protected override void CheckInputSchema(SchemaShape inputSchema)
TextDataViewType.Instance.ToString(), sentenceCol2.GetTypeString());
}
}
else if (BertOptions.TaskType == BertTaskType.NameEntityRecognition)
else if (BertOptions.TaskType == BertTaskType.NamedEntityRecognition)
{
if (labelCol.ItemType != NumberDataViewType.UInt32)
throw Host.ExceptSchemaMismatch(nameof(inputSchema), "label", Option.LabelColumnName,
Expand DownExpand Up@@ -535,7 +535,7 @@ protected override DataViewSchema.DetachedColumn[] GetOutputColumnsCore()
info[1] = new DataViewSchema.DetachedColumn(Parent.Options.ScoreColumnName, new VectorDataViewType(NumberDataViewType.Single, Parent.Options.NumberOfClasses), meta.ToAnnotations());
return info;
}
else if (Parent.BertOptions.TaskType == BertTaskType.NameEntityRecognition)
else if (Parent.BertOptions.TaskType == BertTaskType.NamedEntityRecognition)
{
var info = new DataViewSchema.DetachedColumn[1];
var keyType = Parent.LabelColumn.Annotations.Schema.GetColumnOrNull(AnnotationUtils.Kinds.KeyValues)?.Type as VectorDataViewType;
Expand Down
8 changes: 4 additions & 4 deletions src/Microsoft.ML.TorchSharp/NasBert/NerTrainer.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -35,7 +35,7 @@ namespace Microsoft.ML.TorchSharp.NasBert
/// </summary>
/// <remarks>
/// <format type="text/markdown"><![CDATA[
/// To create this trainer, use [NER](xref:Microsoft.ML.TorchSharpCatalog.NameEntityRecognition(Microsoft.ML.MulticlassClassificationCatalog.MulticlassClassificationTrainers,System.String,System.String,System.String,Int32,Int32,Int32,Microsoft.ML.TorchSharp.NasBert.BertArchitecture,Microsoft.ML.IDataView)).
/// To create this trainer, use [NER](xref:Microsoft.ML.TorchSharpCatalog.NamedEntityRecognition(Microsoft.ML.MulticlassClassificationCatalog.MulticlassClassificationTrainers,System.String,System.String,System.String,Int32,Int32,Int32,Microsoft.ML.TorchSharp.NasBert.BertArchitecture,Microsoft.ML.IDataView)).
///
/// ### Input and Output Columns
/// The input label column data must be a Vector of [string](xref:Microsoft.ML.Data.TextDataViewType) type and the sentence columns must be of type<xref:Microsoft.ML.Data.TextDataViewType>.
Expand All@@ -54,7 +54,7 @@ namespace Microsoft.ML.TorchSharp.NasBert
/// | Exportable to ONNX | No |
///
/// ### Training Algorithm Details
/// Trains a Deep Neural Network(DNN) by leveraging an existing pre-trained NAS-BERT roBERTa model for the purpose of name entity recognition.
/// Trains a Deep Neural Network(DNN) by leveraging an existing pre-trained NAS-BERT roBERTa model for the purpose of named entity recognition.
/// ]]>
/// </format>
/// </remarks>
Expand DownExpand Up@@ -93,7 +93,7 @@ internal NerTrainer(IHostEnvironment env,
BatchSize = batchSize,
MaxEpoch = maxEpochs,
ValidationSet = validationSet,
TaskType = BertTaskType.NameEntityRecognition
TaskType = BertTaskType.NamedEntityRecognition
})
{
}
Expand DownExpand Up@@ -295,7 +295,7 @@ private static NerTransformer Create(IHostEnvironment env, ModelLoadContext ctx)

options.Sentence1ColumnName = ctx.LoadString();
options.Sentence2ColumnName = ctx.LoadStringOrNull();
options.TaskType = BertTaskType.NameEntityRecognition;
options.TaskType = BertTaskType.NamedEntityRecognition;

BinarySaver saver = new BinarySaver(env, new BinarySaver.Arguments());
DataViewType type;
Expand Down
47 changes: 43 additions & 4 deletions src/Microsoft.ML.TorchSharp/TorchSharpCatalog.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -4,6 +4,7 @@

using System;
using System.Collections.Generic;
using System.ComponentModel;
using System.Text;
using Microsoft.ML.Data;
using Microsoft.ML.TorchSharp.AutoFormerV2;
Expand DownExpand Up@@ -161,7 +162,45 @@ public static ObjectDetectionMetrics EvaluateObjectDetection(
}

/// <summary>
/// Fine tune a NAS-BERT model for Name Entity Recognition. The limit for any sentence is 512 tokens. Each word typically
/// Obsolete: please use the <see cref="NamedEntityRecognition(MulticlassClassificationCatalog.MulticlassClassificationTrainers, string, string, string, int, int, BertArchitecture, IDataView)"/> method instead
/// </summary>
/// <param name="catalog">The transform's catalog.</param>
/// <param name="labelColumnName">Name of the label column. Column should be a key type.</param>
/// <param name="outputColumnName">Name of the output column. It will be a key type. It is the predicted label.</param>
/// <param name="sentence1ColumnName">Name of the column for the first sentence.</param>
/// <param name="batchSize">Number of rows in the batch.</param>
/// <param name="maxEpochs">Maximum number of times to loop through your training set.</param>
/// <param name="architecture">Architecture for the model. Defaults to Roberta.</param>
/// <param name="validationSet">The validation set used while training to improve model quality.</param>
/// <returns></returns>
[Obsolete("Please use NamedEntityRecognition method instead", false)]
[EditorBrowsable(EditorBrowsableState.Never)]
public static NerTrainer NameEntityRecognition(
this MulticlassClassificationCatalog.MulticlassClassificationTrainers catalog,
string labelColumnName = DefaultColumnNames.Label,
string outputColumnName = DefaultColumnNames.PredictedLabel,
string sentence1ColumnName = "Sentence",
int batchSize = 32,
int maxEpochs = 10,
BertArchitecture architecture = BertArchitecture.Roberta,
IDataView validationSet = null)
=> NamedEntityRecognition(catalog, labelColumnName, outputColumnName, sentence1ColumnName, batchSize, maxEpochs, architecture, validationSet);

/// <summary>
/// Obsolete: please use the <see cref="NamedEntityRecognition(MulticlassClassificationCatalog.MulticlassClassificationTrainers, NerTrainer.NerOptions)"/> method instead
/// </summary>
/// <param name="catalog">The transform's catalog.</param>
/// <param name="options">The full set of advanced options.</param>
/// <returns></returns>
[Obsolete("Please use NamedEntityRecognition method instead", false)]
[EditorBrowsable(EditorBrowsableState.Never)]
public static NerTrainer NameEntityRecognition(
this MulticlassClassificationCatalog.MulticlassClassificationTrainers catalog,
NerTrainer.NerOptions options)
=> NamedEntityRecognition(catalog, options);

/// <summary>
/// Fine tune a NAS-BERT model for Named Entity Recognition. The limit for any sentence is 512 tokens. Each word typically
/// will map to a single token, and we automatically add 2 specical tokens (a start token and a separator token)
/// so in general this limit will be 510 words for all sentences.
/// </summary>
Expand All@@ -174,7 +213,7 @@ public static ObjectDetectionMetrics EvaluateObjectDetection(
/// <param name="architecture">Architecture for the model. Defaults to Roberta.</param>
/// <param name="validationSet">The validation set used while training to improve model quality.</param>
/// <returns></returns>
public static NerTrainer NameEntityRecognition(
public static NerTrainer NamedEntityRecognition(
this MulticlassClassificationCatalog.MulticlassClassificationTrainers catalog,
string labelColumnName = DefaultColumnNames.Label,
string outputColumnName = DefaultColumnNames.PredictedLabel,
Expand All@@ -186,12 +225,12 @@ public static NerTrainer NameEntityRecognition(
=> new NerTrainer(CatalogUtils.GetEnvironment(catalog), labelColumnName, outputColumnName, sentence1ColumnName, batchSize, maxEpochs, validationSet, architecture);

/// <summary>
/// Fine tune a Name Entity Recognition model.
/// Fine tune a Named Entity Recognition model.
/// </summary>
/// <param name="catalog">The transform's catalog.</param>
/// <param name="options">The full set of advanced options.</param>
/// <returns></returns>
public static NerTrainer NameEntityRecognition(
public static NerTrainer NamedEntityRecognition(
this MulticlassClassificationCatalog.MulticlassClassificationTrainers catalog,
NerTrainer.NerOptions options)
=> new NerTrainer(CatalogUtils.GetEnvironment(catalog), options);
Expand Down
2 changes: 1 addition & 1 deletion test/Microsoft.ML.Tests/NerTests.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -54,7 +54,7 @@ public void TestSimpleNer()
}));
var chain = new EstimatorChain<ITransformer>();
var estimator = chain.Append(ML.Transforms.Conversion.MapValueToKey("Label", keyData: labels))
.Append(ML.MulticlassClassification.Trainers.NameEntityRecognition(outputColumnName: "outputColumn"))
.Append(ML.MulticlassClassification.Trainers.NamedEntityRecognition(outputColumnName: "outputColumn"))
.Append(ML.Transforms.Conversion.MapKeyToValue("outputColumn"));

var estimatorSchema = estimator.GetOutputSchema(SchemaShape.Create(dataView.Schema));
Expand Down
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Remove or un-stick sticky/fixed headers that block content\n(function() {\n function unstick() {\n document.querySelectorAll('header, nav, [role=\"banner\"], .header, .navbar, .sticky, .fixed-top, [style*=\"position: fixed\"], [style*=\"position:sticky\"]').forEach(function(el) {\n if (el.style.position === 'fixed' || el.style.position === 'sticky' || \n getComputedStyle(el).position === 'fixed' || getComputedStyle(el).position === 'sticky') {\n el.style.position = 'static';\n el.style.top = 'auto';\n el.style.zIndex = 'auto';\n }\n });\n }\n \n unstick();\n \n var observer = new MutationObserver(unstick);\n observer.observe(document.body, { childList: true, subtree: true, attributes: true, attributeFilter: ['style', 'class'] });\n})();", "Kill Sticky Headers"); } } catch(__e) { console.warn('[Userscript:Kill Sticky Headers]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 5 additions & 1 deletion src/Microsoft.ML.TorchSharp/NasBert/BertTaskType.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -4,6 +4,7 @@

using System;
using System.Collections.Generic;
using System.ComponentModel;
using System.Text;

namespace Microsoft.ML.TorchSharp.NasBert
Expand All@@ -17,7 +18,10 @@ public enum BertTaskType
MaskedLM = 1,
TextClassification = 2,
SentenceRegression = 3,
NameEntityRecognition = 4,
NamedEntityRecognition = 4,
[Obsolete("Please use NamedEntityRecognition instead", false)]
[EditorBrowsable(EditorBrowsableState.Never)]
NameEntityRecognition = NamedEntityRecognition,
QuestionAnswering = 5
}
}
12 changes: 6 additions & 6 deletions src/Microsoft.ML.TorchSharp/NasBert/NasBertTrainer.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -204,7 +204,7 @@ private protected override Module CreateModule(IChannel ch, IDataView input)
EnglishRoberta tokenizerModel = Tokenizer.RobertaModel();

NasBertModel model;
if (Parent.BertOptions.TaskType == BertTaskType.NameEntityRecognition)
if (Parent.BertOptions.TaskType == BertTaskType.NamedEntityRecognition)
model = new NerModel(Parent.BertOptions, tokenizerModel.PadIndex, tokenizerModel.SymbolsCount, Parent.Option.NumberOfClasses);
else
model = new ModelForPrediction(Parent.BertOptions, tokenizerModel.PadIndex, tokenizerModel.SymbolsCount, Parent.Option.NumberOfClasses);
Expand DownExpand Up@@ -268,7 +268,7 @@ private protected override torch.Tensor PrepareRowTensor()
private protected override void RunModelAndBackPropagate(ref List<Tensor> inputTensors, ref Tensor targetsTensor)
{
Tensor logits = default;
if (Parent.BertOptions.TaskType == BertTaskType.NameEntityRecognition)
if (Parent.BertOptions.TaskType == BertTaskType.NamedEntityRecognition)
{
int[,] lengthArray = new int[inputTensors.Count, 1];
for (int i = 0; i < inputTensors.Count; i++)
Expand All@@ -293,7 +293,7 @@ private protected override void RunModelAndBackPropagate(ref List<Tensor> inputT
torch.Tensor loss;
if (Parent.BertOptions.TaskType == BertTaskType.TextClassification)
loss = torch.nn.CrossEntropyLoss(reduction: Parent.BertOptions.Reduction).forward(logits, targetsTensor);
else if (Parent.BertOptions.TaskType == BertTaskType.NameEntityRecognition)
else if (Parent.BertOptions.TaskType == BertTaskType.NamedEntityRecognition)
{
targetsTensor = targetsTensor.@long().view(-1);
logits = logits.view(-1, logits.size(-1));
Expand DownExpand Up@@ -338,7 +338,7 @@ public override SchemaShape GetOutputSchema(SchemaShape inputSchema)
outColumns[Option.ScoreColumnName] = new SchemaShape.Column(Option.ScoreColumnName, SchemaShape.Column.VectorKind.Vector,
NumberDataViewType.Single, false, new SchemaShape(AnnotationUtils.AnnotationsForMulticlassScoreColumn(labelCol)));
}
else if (BertOptions.TaskType == BertTaskType.NameEntityRecognition)
else if (BertOptions.TaskType == BertTaskType.NamedEntityRecognition)
{
var metadata = new List<SchemaShape.Column>();
metadata.Add(new SchemaShape.Column(AnnotationUtils.Kinds.KeyValues, SchemaShape.Column.VectorKind.Vector,
Expand DownExpand Up@@ -387,7 +387,7 @@ private protected override void CheckInputSchema(SchemaShape inputSchema)
TextDataViewType.Instance.ToString(), sentenceCol2.GetTypeString());
}
}
else if (BertOptions.TaskType == BertTaskType.NameEntityRecognition)
else if (BertOptions.TaskType == BertTaskType.NamedEntityRecognition)
{
if (labelCol.ItemType != NumberDataViewType.UInt32)
throw Host.ExceptSchemaMismatch(nameof(inputSchema), "label", Option.LabelColumnName,
Expand DownExpand Up@@ -535,7 +535,7 @@ protected override DataViewSchema.DetachedColumn[] GetOutputColumnsCore()
info[1] = new DataViewSchema.DetachedColumn(Parent.Options.ScoreColumnName, new VectorDataViewType(NumberDataViewType.Single, Parent.Options.NumberOfClasses), meta.ToAnnotations());
return info;
}
else if (Parent.BertOptions.TaskType == BertTaskType.NameEntityRecognition)
else if (Parent.BertOptions.TaskType == BertTaskType.NamedEntityRecognition)
{
var info = new DataViewSchema.DetachedColumn[1];
var keyType = Parent.LabelColumn.Annotations.Schema.GetColumnOrNull(AnnotationUtils.Kinds.KeyValues)?.Type as VectorDataViewType;
Expand Down
8 changes: 4 additions & 4 deletions src/Microsoft.ML.TorchSharp/NasBert/NerTrainer.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -35,7 +35,7 @@ namespace Microsoft.ML.TorchSharp.NasBert
/// </summary>
/// <remarks>
/// <format type="text/markdown"><![CDATA[
/// To create this trainer, use [NER](xref:Microsoft.ML.TorchSharpCatalog.NameEntityRecognition(Microsoft.ML.MulticlassClassificationCatalog.MulticlassClassificationTrainers,System.String,System.String,System.String,Int32,Int32,Int32,Microsoft.ML.TorchSharp.NasBert.BertArchitecture,Microsoft.ML.IDataView)).
/// To create this trainer, use [NER](xref:Microsoft.ML.TorchSharpCatalog.NamedEntityRecognition(Microsoft.ML.MulticlassClassificationCatalog.MulticlassClassificationTrainers,System.String,System.String,System.String,Int32,Int32,Int32,Microsoft.ML.TorchSharp.NasBert.BertArchitecture,Microsoft.ML.IDataView)).
///
/// ### Input and Output Columns
/// The input label column data must be a Vector of [string](xref:Microsoft.ML.Data.TextDataViewType) type and the sentence columns must be of type<xref:Microsoft.ML.Data.TextDataViewType>.
Expand All@@ -54,7 +54,7 @@ namespace Microsoft.ML.TorchSharp.NasBert
/// | Exportable to ONNX | No |
///
/// ### Training Algorithm Details
/// Trains a Deep Neural Network(DNN) by leveraging an existing pre-trained NAS-BERT roBERTa model for the purpose of name entity recognition.
/// Trains a Deep Neural Network(DNN) by leveraging an existing pre-trained NAS-BERT roBERTa model for the purpose of named entity recognition.
/// ]]>
/// </format>
/// </remarks>
Expand DownExpand Up@@ -93,7 +93,7 @@ internal NerTrainer(IHostEnvironment env,
BatchSize = batchSize,
MaxEpoch = maxEpochs,
ValidationSet = validationSet,
TaskType = BertTaskType.NameEntityRecognition
TaskType = BertTaskType.NamedEntityRecognition
})
{
}
Expand DownExpand Up@@ -295,7 +295,7 @@ private static NerTransformer Create(IHostEnvironment env, ModelLoadContext ctx)

options.Sentence1ColumnName = ctx.LoadString();
options.Sentence2ColumnName = ctx.LoadStringOrNull();
options.TaskType = BertTaskType.NameEntityRecognition;
options.TaskType = BertTaskType.NamedEntityRecognition;

BinarySaver saver = new BinarySaver(env, new BinarySaver.Arguments());
DataViewType type;
Expand Down
47 changes: 43 additions & 4 deletions src/Microsoft.ML.TorchSharp/TorchSharpCatalog.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -4,6 +4,7 @@

using System;
using System.Collections.Generic;
using System.ComponentModel;
using System.Text;
using Microsoft.ML.Data;
using Microsoft.ML.TorchSharp.AutoFormerV2;
Expand DownExpand Up@@ -161,7 +162,45 @@ public static ObjectDetectionMetrics EvaluateObjectDetection(
}

/// <summary>
/// Fine tune a NAS-BERT model for Name Entity Recognition. The limit for any sentence is 512 tokens. Each word typically
/// Obsolete: please use the <see cref="NamedEntityRecognition(MulticlassClassificationCatalog.MulticlassClassificationTrainers, string, string, string, int, int, BertArchitecture, IDataView)"/> method instead
/// </summary>
/// <param name="catalog">The transform's catalog.</param>
/// <param name="labelColumnName">Name of the label column. Column should be a key type.</param>
/// <param name="outputColumnName">Name of the output column. It will be a key type. It is the predicted label.</param>
/// <param name="sentence1ColumnName">Name of the column for the first sentence.</param>
/// <param name="batchSize">Number of rows in the batch.</param>
/// <param name="maxEpochs">Maximum number of times to loop through your training set.</param>
/// <param name="architecture">Architecture for the model. Defaults to Roberta.</param>
/// <param name="validationSet">The validation set used while training to improve model quality.</param>
/// <returns></returns>
[Obsolete("Please use NamedEntityRecognition method instead", false)]
[EditorBrowsable(EditorBrowsableState.Never)]
public static NerTrainer NameEntityRecognition(
this MulticlassClassificationCatalog.MulticlassClassificationTrainers catalog,
string labelColumnName = DefaultColumnNames.Label,
string outputColumnName = DefaultColumnNames.PredictedLabel,
string sentence1ColumnName = "Sentence",
int batchSize = 32,
int maxEpochs = 10,
BertArchitecture architecture = BertArchitecture.Roberta,
IDataView validationSet = null)
=> NamedEntityRecognition(catalog, labelColumnName, outputColumnName, sentence1ColumnName, batchSize, maxEpochs, architecture, validationSet);

/// <summary>
/// Obsolete: please use the <see cref="NamedEntityRecognition(MulticlassClassificationCatalog.MulticlassClassificationTrainers, NerTrainer.NerOptions)"/> method instead
/// </summary>
/// <param name="catalog">The transform's catalog.</param>
/// <param name="options">The full set of advanced options.</param>
/// <returns></returns>
[Obsolete("Please use NamedEntityRecognition method instead", false)]
[EditorBrowsable(EditorBrowsableState.Never)]
public static NerTrainer NameEntityRecognition(
this MulticlassClassificationCatalog.MulticlassClassificationTrainers catalog,
NerTrainer.NerOptions options)
=> NamedEntityRecognition(catalog, options);

/// <summary>
/// Fine tune a NAS-BERT model for Named Entity Recognition. The limit for any sentence is 512 tokens. Each word typically
/// will map to a single token, and we automatically add 2 specical tokens (a start token and a separator token)
/// so in general this limit will be 510 words for all sentences.
/// </summary>
Expand All@@ -174,7 +213,7 @@ public static ObjectDetectionMetrics EvaluateObjectDetection(
/// <param name="architecture">Architecture for the model. Defaults to Roberta.</param>
/// <param name="validationSet">The validation set used while training to improve model quality.</param>
/// <returns></returns>
public static NerTrainer NameEntityRecognition(
public static NerTrainer NamedEntityRecognition(
this MulticlassClassificationCatalog.MulticlassClassificationTrainers catalog,
string labelColumnName = DefaultColumnNames.Label,
string outputColumnName = DefaultColumnNames.PredictedLabel,
Expand All@@ -186,12 +225,12 @@ public static NerTrainer NameEntityRecognition(
=> new NerTrainer(CatalogUtils.GetEnvironment(catalog), labelColumnName, outputColumnName, sentence1ColumnName, batchSize, maxEpochs, validationSet, architecture);

/// <summary>
/// Fine tune a Name Entity Recognition model.
/// Fine tune a Named Entity Recognition model.
/// </summary>
/// <param name="catalog">The transform's catalog.</param>
/// <param name="options">The full set of advanced options.</param>
/// <returns></returns>
public static NerTrainer NameEntityRecognition(
public static NerTrainer NamedEntityRecognition(
this MulticlassClassificationCatalog.MulticlassClassificationTrainers catalog,
NerTrainer.NerOptions options)
=> new NerTrainer(CatalogUtils.GetEnvironment(catalog), options);
Expand Down
2 changes: 1 addition & 1 deletion test/Microsoft.ML.Tests/NerTests.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -54,7 +54,7 @@ public void TestSimpleNer()
}));
var chain = new EstimatorChain<ITransformer>();
var estimator = chain.Append(ML.Transforms.Conversion.MapValueToKey("Label", keyData: labels))
.Append(ML.MulticlassClassification.Trainers.NameEntityRecognition(outputColumnName: "outputColumn"))
.Append(ML.MulticlassClassification.Trainers.NamedEntityRecognition(outputColumnName: "outputColumn"))
.Append(ML.Transforms.Conversion.MapKeyToValue("outputColumn"));

var estimatorSchema = estimator.GetOutputSchema(SchemaShape.Create(dataView.Schema));
Expand Down
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Universal Dark Mode - works on any site\n(function() {\n var enabled = true;\n \n function applyDarkMode() {\n if (!enabled) return;\n \n // Create style element if it doesn't exist\n var style = document.getElementById('universal-dark-mode-style');\n if (!style) {\n style = document.createElement('style');\n style.id = 'universal-dark-mode-style';\n document.head.appendChild(style);\n }\n \n // Dark mode CSS - inverts colors but preserves images/video\n style.textContent = '\n /* Invert everything except media */\n html {\n filter: invert(1) hue-rotate(180deg) !important;\n background: #1a1a2e !important;\n }\n \n /* Restore images, videos, iframes, canvas */\n img, video, iframe, canvas, svg, picture, [style*=\"background-image\"] {\n filter: invert(1) hue-rotate(180deg) !important;\n }\n \n /* Preserve specific elements that should not be inverted */\n .no-dark-mode, .no-dark-mode *,\n [data-theme=\"light\"], [data-theme=\"light\"],\n .ace_editor, .ace_editor *,\n .CodeMirror, .CodeMirror *,\n .monaco-editor, .monaco-editor *,\n .markdown-body pre, .markdown-body pre *,\n .highlight, .highlight *,\n pre code, pre code * {\n filter: none !important;\n }\n \n /* Fix common UI elements */\n .modal, .popup, .dropdown-menu, .tooltip, .popover {\n filter: invert(1) hue-rotate(180deg) !important;\n background: #2d2d44 !important;\n border-color: #444 !important;\n }\n \n /* Scrollbars */\n ::-webkit-scrollbar { background: #1a1a2e !important; }\n ::-webkit-scrollbar-thumb { background: #444 !important; }\n ::-webkit-scrollbar-thumb:hover { background: #555 !important; }\n \n /* Selection */\n ::selection { background: #4ecdc4 !important; color: #1a1a2e !important; }\n ::-moz-selection { background: #4ecdc4 !important; color: #1a1a2e !important; }\n ';\n }\n \n function removeDarkMode() {\n var style = document.getElementById('universal-dark-mode-style');\n if (style) style.remove();\n }\n \n // Toggle with Alt+Shift+D\n document.addEventListener('keydown', function(e) {\n if (e.altKey && e.shiftKey && e.key === 'D') {\n e.preventDefault();\n enabled = !enabled;\n if (enabled) {\n applyDarkMode();\n console.log('[Universal Dark Mode] Enabled');\n } else {\n removeDarkMode();\n console.log('[Universal Dark Mode] Disabled');\n }\n }\n });\n \n // Apply on load\n applyDarkMode();\n \n // Re-apply on dynamic content\n var observer = new MutationObserver(function(mutations) {\n if (enabled && !document.getElementById('universal-dark-mode-style')) {\n applyDarkMode();\n }\n });\n observer.observe(document.head, { childList: true });\n \n console.log('[Universal Dark Mode] Loaded - Press Alt+Shift+D to toggle');\n})();", "Universal Dark Mode"); } } catch(__e) { console.warn('[Userscript:Universal Dark Mode]', __e); } })(); })();
Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 5 additions & 1 deletion src/Microsoft.ML.TorchSharp/NasBert/BertTaskType.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -4,6 +4,7 @@

using System;
using System.Collections.Generic;
using System.ComponentModel;
using System.Text;

namespace Microsoft.ML.TorchSharp.NasBert
Expand All@@ -17,7 +18,10 @@ public enum BertTaskType
MaskedLM = 1,
TextClassification = 2,
SentenceRegression = 3,
NameEntityRecognition = 4,
NamedEntityRecognition = 4,
[Obsolete("Please use NamedEntityRecognition instead", false)]
[EditorBrowsable(EditorBrowsableState.Never)]
NameEntityRecognition = NamedEntityRecognition,
QuestionAnswering = 5
}
}
12 changes: 6 additions & 6 deletions src/Microsoft.ML.TorchSharp/NasBert/NasBertTrainer.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -204,7 +204,7 @@ private protected override Module CreateModule(IChannel ch, IDataView input)
EnglishRoberta tokenizerModel = Tokenizer.RobertaModel();

NasBertModel model;
if (Parent.BertOptions.TaskType == BertTaskType.NameEntityRecognition)
if (Parent.BertOptions.TaskType == BertTaskType.NamedEntityRecognition)
model = new NerModel(Parent.BertOptions, tokenizerModel.PadIndex, tokenizerModel.SymbolsCount, Parent.Option.NumberOfClasses);
else
model = new ModelForPrediction(Parent.BertOptions, tokenizerModel.PadIndex, tokenizerModel.SymbolsCount, Parent.Option.NumberOfClasses);
Expand DownExpand Up@@ -268,7 +268,7 @@ private protected override torch.Tensor PrepareRowTensor()
private protected override void RunModelAndBackPropagate(ref List<Tensor> inputTensors, ref Tensor targetsTensor)
{
Tensor logits = default;
if (Parent.BertOptions.TaskType == BertTaskType.NameEntityRecognition)
if (Parent.BertOptions.TaskType == BertTaskType.NamedEntityRecognition)
{
int[,] lengthArray = new int[inputTensors.Count, 1];
for (int i = 0; i < inputTensors.Count; i++)
Expand All@@ -293,7 +293,7 @@ private protected override void RunModelAndBackPropagate(ref List<Tensor> inputT
torch.Tensor loss;
if (Parent.BertOptions.TaskType == BertTaskType.TextClassification)
loss = torch.nn.CrossEntropyLoss(reduction: Parent.BertOptions.Reduction).forward(logits, targetsTensor);
else if (Parent.BertOptions.TaskType == BertTaskType.NameEntityRecognition)
else if (Parent.BertOptions.TaskType == BertTaskType.NamedEntityRecognition)
{
targetsTensor = targetsTensor.@long().view(-1);
logits = logits.view(-1, logits.size(-1));
Expand DownExpand Up@@ -338,7 +338,7 @@ public override SchemaShape GetOutputSchema(SchemaShape inputSchema)
outColumns[Option.ScoreColumnName] = new SchemaShape.Column(Option.ScoreColumnName, SchemaShape.Column.VectorKind.Vector,
NumberDataViewType.Single, false, new SchemaShape(AnnotationUtils.AnnotationsForMulticlassScoreColumn(labelCol)));
}
else if (BertOptions.TaskType == BertTaskType.NameEntityRecognition)
else if (BertOptions.TaskType == BertTaskType.NamedEntityRecognition)
{
var metadata = new List<SchemaShape.Column>();
metadata.Add(new SchemaShape.Column(AnnotationUtils.Kinds.KeyValues, SchemaShape.Column.VectorKind.Vector,
Expand DownExpand Up@@ -387,7 +387,7 @@ private protected override void CheckInputSchema(SchemaShape inputSchema)
TextDataViewType.Instance.ToString(), sentenceCol2.GetTypeString());
}
}
else if (BertOptions.TaskType == BertTaskType.NameEntityRecognition)
else if (BertOptions.TaskType == BertTaskType.NamedEntityRecognition)
{
if (labelCol.ItemType != NumberDataViewType.UInt32)
throw Host.ExceptSchemaMismatch(nameof(inputSchema), "label", Option.LabelColumnName,
Expand DownExpand Up@@ -535,7 +535,7 @@ protected override DataViewSchema.DetachedColumn[] GetOutputColumnsCore()
info[1] = new DataViewSchema.DetachedColumn(Parent.Options.ScoreColumnName, new VectorDataViewType(NumberDataViewType.Single, Parent.Options.NumberOfClasses), meta.ToAnnotations());
return info;
}
else if (Parent.BertOptions.TaskType == BertTaskType.NameEntityRecognition)
else if (Parent.BertOptions.TaskType == BertTaskType.NamedEntityRecognition)
{
var info = new DataViewSchema.DetachedColumn[1];
var keyType = Parent.LabelColumn.Annotations.Schema.GetColumnOrNull(AnnotationUtils.Kinds.KeyValues)?.Type as VectorDataViewType;
Expand Down
8 changes: 4 additions & 4 deletions src/Microsoft.ML.TorchSharp/NasBert/NerTrainer.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -35,7 +35,7 @@ namespace Microsoft.ML.TorchSharp.NasBert
/// </summary>
/// <remarks>
/// <format type="text/markdown"><![CDATA[
/// To create this trainer, use [NER](xref:Microsoft.ML.TorchSharpCatalog.NameEntityRecognition(Microsoft.ML.MulticlassClassificationCatalog.MulticlassClassificationTrainers,System.String,System.String,System.String,Int32,Int32,Int32,Microsoft.ML.TorchSharp.NasBert.BertArchitecture,Microsoft.ML.IDataView)).
/// To create this trainer, use [NER](xref:Microsoft.ML.TorchSharpCatalog.NamedEntityRecognition(Microsoft.ML.MulticlassClassificationCatalog.MulticlassClassificationTrainers,System.String,System.String,System.String,Int32,Int32,Int32,Microsoft.ML.TorchSharp.NasBert.BertArchitecture,Microsoft.ML.IDataView)).
///
/// ### Input and Output Columns
/// The input label column data must be a Vector of [string](xref:Microsoft.ML.Data.TextDataViewType) type and the sentence columns must be of type<xref:Microsoft.ML.Data.TextDataViewType>.
Expand All@@ -54,7 +54,7 @@ namespace Microsoft.ML.TorchSharp.NasBert
/// | Exportable to ONNX | No |
///
/// ### Training Algorithm Details
/// Trains a Deep Neural Network(DNN) by leveraging an existing pre-trained NAS-BERT roBERTa model for the purpose of name entity recognition.
/// Trains a Deep Neural Network(DNN) by leveraging an existing pre-trained NAS-BERT roBERTa model for the purpose of named entity recognition.
/// ]]>
/// </format>
/// </remarks>
Expand DownExpand Up@@ -93,7 +93,7 @@ internal NerTrainer(IHostEnvironment env,
BatchSize = batchSize,
MaxEpoch = maxEpochs,
ValidationSet = validationSet,
TaskType = BertTaskType.NameEntityRecognition
TaskType = BertTaskType.NamedEntityRecognition
})
{
}
Expand DownExpand Up@@ -295,7 +295,7 @@ private static NerTransformer Create(IHostEnvironment env, ModelLoadContext ctx)

options.Sentence1ColumnName = ctx.LoadString();
options.Sentence2ColumnName = ctx.LoadStringOrNull();
options.TaskType = BertTaskType.NameEntityRecognition;
options.TaskType = BertTaskType.NamedEntityRecognition;

BinarySaver saver = new BinarySaver(env, new BinarySaver.Arguments());
DataViewType type;
Expand Down
47 changes: 43 additions & 4 deletions src/Microsoft.ML.TorchSharp/TorchSharpCatalog.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -4,6 +4,7 @@

using System;
using System.Collections.Generic;
using System.ComponentModel;
using System.Text;
using Microsoft.ML.Data;
using Microsoft.ML.TorchSharp.AutoFormerV2;
Expand DownExpand Up@@ -161,7 +162,45 @@ public static ObjectDetectionMetrics EvaluateObjectDetection(
}

/// <summary>
/// Fine tune a NAS-BERT model for Name Entity Recognition. The limit for any sentence is 512 tokens. Each word typically
/// Obsolete: please use the <see cref="NamedEntityRecognition(MulticlassClassificationCatalog.MulticlassClassificationTrainers, string, string, string, int, int, BertArchitecture, IDataView)"/> method instead
/// </summary>
/// <param name="catalog">The transform's catalog.</param>
/// <param name="labelColumnName">Name of the label column. Column should be a key type.</param>
/// <param name="outputColumnName">Name of the output column. It will be a key type. It is the predicted label.</param>
/// <param name="sentence1ColumnName">Name of the column for the first sentence.</param>
/// <param name="batchSize">Number of rows in the batch.</param>
/// <param name="maxEpochs">Maximum number of times to loop through your training set.</param>
/// <param name="architecture">Architecture for the model. Defaults to Roberta.</param>
/// <param name="validationSet">The validation set used while training to improve model quality.</param>
/// <returns></returns>
[Obsolete("Please use NamedEntityRecognition method instead", false)]
[EditorBrowsable(EditorBrowsableState.Never)]
public static NerTrainer NameEntityRecognition(
this MulticlassClassificationCatalog.MulticlassClassificationTrainers catalog,
string labelColumnName = DefaultColumnNames.Label,
string outputColumnName = DefaultColumnNames.PredictedLabel,
string sentence1ColumnName = "Sentence",
int batchSize = 32,
int maxEpochs = 10,
BertArchitecture architecture = BertArchitecture.Roberta,
IDataView validationSet = null)
=> NamedEntityRecognition(catalog, labelColumnName, outputColumnName, sentence1ColumnName, batchSize, maxEpochs, architecture, validationSet);

/// <summary>
/// Obsolete: please use the <see cref="NamedEntityRecognition(MulticlassClassificationCatalog.MulticlassClassificationTrainers, NerTrainer.NerOptions)"/> method instead
/// </summary>
/// <param name="catalog">The transform's catalog.</param>
/// <param name="options">The full set of advanced options.</param>
/// <returns></returns>
[Obsolete("Please use NamedEntityRecognition method instead", false)]
[EditorBrowsable(EditorBrowsableState.Never)]
public static NerTrainer NameEntityRecognition(
this MulticlassClassificationCatalog.MulticlassClassificationTrainers catalog,
NerTrainer.NerOptions options)
=> NamedEntityRecognition(catalog, options);

/// <summary>
/// Fine tune a NAS-BERT model for Named Entity Recognition. The limit for any sentence is 512 tokens. Each word typically
/// will map to a single token, and we automatically add 2 specical tokens (a start token and a separator token)
/// so in general this limit will be 510 words for all sentences.
/// </summary>
Expand All@@ -174,7 +213,7 @@ public static ObjectDetectionMetrics EvaluateObjectDetection(
/// <param name="architecture">Architecture for the model. Defaults to Roberta.</param>
/// <param name="validationSet">The validation set used while training to improve model quality.</param>
/// <returns></returns>
public static NerTrainer NameEntityRecognition(
public static NerTrainer NamedEntityRecognition(
this MulticlassClassificationCatalog.MulticlassClassificationTrainers catalog,
string labelColumnName = DefaultColumnNames.Label,
string outputColumnName = DefaultColumnNames.PredictedLabel,
Expand All@@ -186,12 +225,12 @@ public static NerTrainer NameEntityRecognition(
=> new NerTrainer(CatalogUtils.GetEnvironment(catalog), labelColumnName, outputColumnName, sentence1ColumnName, batchSize, maxEpochs, validationSet, architecture);

/// <summary>
/// Fine tune a Name Entity Recognition model.
/// Fine tune a Named Entity Recognition model.
/// </summary>
/// <param name="catalog">The transform's catalog.</param>
/// <param name="options">The full set of advanced options.</param>
/// <returns></returns>
public static NerTrainer NameEntityRecognition(
public static NerTrainer NamedEntityRecognition(
this MulticlassClassificationCatalog.MulticlassClassificationTrainers catalog,
NerTrainer.NerOptions options)
=> new NerTrainer(CatalogUtils.GetEnvironment(catalog), options);
Expand Down
2 changes: 1 addition & 1 deletion test/Microsoft.ML.Tests/NerTests.cs
Original file line numberDiff line numberDiff line change
Expand Up@@ -54,7 +54,7 @@ public void TestSimpleNer()
}));
var chain = new EstimatorChain<ITransformer>();
var estimator = chain.Append(ML.Transforms.Conversion.MapValueToKey("Label", keyData: labels))
.Append(ML.MulticlassClassification.Trainers.NameEntityRecognition(outputColumnName: "outputColumn"))
.Append(ML.MulticlassClassification.Trainers.NamedEntityRecognition(outputColumnName: "outputColumn"))
.Append(ML.Transforms.Conversion.MapKeyToValue("outputColumn"));

var estimatorSchema = estimator.GetOutputSchema(SchemaShape.Create(dataView.Schema));
Expand Down