diff --git a/src/Microsoft.Data.Analysis/DataFrame.cs b/src/Microsoft.Data.Analysis/DataFrame.cs index c9de526013..3a524d73a8 100644 --- a/src/Microsoft.Data.Analysis/DataFrame.cs +++ b/src/Microsoft.Data.Analysis/DataFrame.cs @@ -177,12 +177,15 @@ public DataFrame Head(int numberOfRows) /// public DataFrame Tail(int numberOfRows) { - PrimitiveDataFrameColumn filter = new PrimitiveDataFrameColumn("Filter", numberOfRows); - for (long i = Rows.Count - numberOfRows; i < Rows.Count; i++) + int rowCount = (int)Rows.Count; + if (numberOfRows >= rowCount) { - filter[i - (Rows.Count - numberOfRows)] = i; + return Clone(new PrimitiveDataFrameColumn("Filter", Enumerable.Range(0, rowCount))); + } + else + { + return Clone(new PrimitiveDataFrameColumn("Filter", Enumerable.Range(rowCount - numberOfRows, numberOfRows))); } - return Clone(filter); } // TODO: Add strongly typed versions of these APIs #endregion @@ -738,99 +741,60 @@ public DataFrame Append(IEnumerable> row, bool inPl /// public DataFrame Melt(IEnumerable idColumns, IEnumerable valueColumns = null, string variableName = "variable", string valueName = "value", bool dropNulls = false) { - if (string.IsNullOrWhiteSpace(variableName)) - { - throw new ArgumentException(Strings.ParameterMustNotBeNullOrWhitespace, nameof(variableName)); - } - - if (string.IsNullOrWhiteSpace(valueName)) - { - throw new ArgumentException(Strings.ParameterMustNotBeNullOrWhitespace, nameof(valueName)); - } - + // Validate input if (idColumns == null) { throw new ArgumentNullException(nameof(idColumns)); } var idColumnList = idColumns.ToList(); + var valueColumnList = valueColumns?.ToList() ?? Columns.Select(c => c.Name).Except(idColumnList).ToList(); - HashSet idColumnSet = null; - - if (valueColumns is null) - { - idColumnSet = [.. idColumnList]; - } - - var valueColumnList = valueColumns?.ToList() - ?? _columnCollection - .Where(c => !idColumnSet.Contains(c.Name)) - .Select(c => c.Name) - .ToList(); - - if (idColumnList.Count == 0) - { - throw new ArgumentException(Strings.MissingIdColumns, nameof(idColumns)); - } - - if (valueColumns != null && valueColumnList.Count == 0) - { - throw new ArgumentException(Strings.MissingValueColumns, nameof(valueColumns)); - } - - if (valueColumns != null && valueColumnList.Any(v => idColumnList.Contains(v))) - { - throw new ArgumentException(Strings.DuplicateColumnsInIdAndValueLists, nameof(valueColumns)); - } - - if (valueColumns == null && valueColumnList.Count == 0) - { - throw new InvalidOperationException(Strings.NoValueColumnsRemaining); - } - - if (_columnCollection.IndexOf(variableName) >= 0) - { - throw new ArgumentException(string.Format(Strings.VariableNameAlreadyExists, variableName), nameof(variableName)); - } - - if (_columnCollection.IndexOf(valueName) >= 0) + // Prepare output columns + var outputColumns = new List(); + foreach (var idCol in idColumnList) { - throw new ArgumentException(string.Format(Strings.ValueNameAlreadyExists, valueName), nameof(valueName)); + outputColumns.Add(Columns[idCol].Clone()); } - - if (string.Equals(variableName, valueName)) + var variableColumn = new StringDataFrameColumn(variableName); + var valueColumn = new DataFrameColumn[valueColumnList.Count]; + for (int i = 0; i < valueColumnList.Count; i++) { - throw new ArgumentException(string.Format(Strings.VariableNameAndValueNameMustBeDifferent, nameof(variableName), nameof(valueName)), nameof(valueName)); + valueColumn[i] = Columns[valueColumnList[i]].Clone(); } + var valueResultColumn = new List(); - foreach (var columnName in idColumnList) + // Helper to add a row to the result + void AddRow(long rowIndex, int valueColIdx) { - if (_columnCollection.IndexOf(columnName) < 0) + foreach (var idCol in idColumnList) { - throw new ArgumentException(string.Format(Strings.InvalidColumnName, columnName), nameof(idColumns)); + outputColumns[idColumnList.IndexOf(idCol)].Append(Columns[idCol][rowIndex]); } + variableColumn.Append(valueColumnList[valueColIdx]); + valueResultColumn.Add(Columns[valueColumnList[valueColIdx]][rowIndex]); } - foreach (var columnName in valueColumnList) + // Main loop + for (long rowIndex = 0; rowIndex < Rows.Count; rowIndex++) { - if (_columnCollection.IndexOf(columnName) < 0) + for (int valueColIdx = 0; valueColIdx < valueColumnList.Count; valueColIdx++) { - throw new ArgumentException(string.Format(Strings.InvalidColumnName, columnName), nameof(valueColumns)); + var value = Columns[valueColumnList[valueColIdx]][rowIndex]; + if (dropNulls && value == null) + { + continue; + } + AddRow(rowIndex, valueColIdx); } } - long totalOutputRows = CalculateTotalOutputRows(valueColumnList, dropNulls); + // Build value column + DataFrameColumn finalValueColumn = Columns[valueColumnList[0]].Clone(valueName, valueResultColumn); - var outputCols = InitializeIdColumns(idColumnList, totalOutputRows); - var variableColumn = new StringDataFrameColumn(variableName, totalOutputRows); - var valueColumn = CreateValueColumn(valueColumnList, valueName, totalOutputRows); - - FillMeltedData(idColumnList, valueColumnList, outputCols, variableColumn, valueColumn, dropNulls); - - outputCols.Add(variableColumn); - outputCols.Add(valueColumn); - - return new DataFrame(outputCols); + // Build result DataFrame + var resultColumns = new List(outputColumns) { variableColumn, finalValueColumn }; + return new DataFrame(resultColumns); } private long CalculateTotalOutputRows(List valueColumnList, bool dropNulls) @@ -1005,5 +969,46 @@ public string ToString(long rowsToShow) return sb.ToString(); } + + /// + /// Determines whether two DataFrame instances are not equal. + /// + public static bool operator !=(DataFrame left, DataFrame right) + { + return !(left == right); + } + + /// + public override bool Equals(object obj) + { + if (obj is DataFrame other) + { + // Compare columns and rows for equality + if (Columns.Count != other.Columns.Count || Rows.Count != other.Rows.Count) + return false; + for (int i = 0; i < Columns.Count; i++) + { + if (!Columns[i].Equals(other.Columns[i])) + return false; + } + // Optionally, compare row data if needed + return true; + } + return false; + } + + /// + public override int GetHashCode() + { + // Combine hash codes of columns and row count + int hash = 17; + hash = hash * 23 + Columns.Count.GetHashCode(); + hash = hash * 23 + Rows.Count.GetHashCode(); + foreach (var column in Columns) + { + hash = hash * 23 + column.GetHashCode(); + } + return hash; + } } }