Updated the normalize command to correct common errors in WhisperX-generated input transcripts
All checks were successful
ci/woodpecker/tag/release Pipeline was successful
All checks were successful
ci/woodpecker/tag/release Pipeline was successful
This commit is contained in:
@@ -19,14 +19,20 @@ const (
|
||||
|
||||
// ParsedTranscript is the validated normalize input model.
|
||||
type ParsedTranscript struct {
|
||||
Shape InputShape
|
||||
Segments []InputSegment
|
||||
Shape InputShape
|
||||
InputSegmentCount int
|
||||
Stats NormalizeStats
|
||||
Segments []InputSegment
|
||||
}
|
||||
|
||||
// InputSegment is a validated segment from normalize input.
|
||||
type InputSegment struct {
|
||||
InputIndex int
|
||||
OriginalID *int
|
||||
StartPresent bool
|
||||
EndPresent bool
|
||||
SpeakerPresent bool
|
||||
TextPresent bool
|
||||
Start float64
|
||||
End float64
|
||||
Speaker string
|
||||
@@ -39,6 +45,14 @@ type InputSegment struct {
|
||||
OverlapGroupID *int
|
||||
}
|
||||
|
||||
// NormalizeStats captures deterministic repair/drop outcomes from normalize input processing.
|
||||
type NormalizeStats struct {
|
||||
TimingFieldsRepaired int
|
||||
TimingOrderSwapped int
|
||||
SpeakerFilled int
|
||||
SegmentsDroppedText int
|
||||
}
|
||||
|
||||
type inputSegmentPayload struct {
|
||||
ID *int `json:"id"`
|
||||
Start *float64 `json:"start"`
|
||||
@@ -85,13 +99,15 @@ func ParseReader(reader io.Reader) (ParsedTranscript, error) {
|
||||
case '{':
|
||||
return parseObjectShape(trimmed)
|
||||
case '[':
|
||||
segments, err := parseSegmentsArray(trimmed)
|
||||
segments, stats, inputCount, err := parseSegmentsArray(trimmed)
|
||||
if err != nil {
|
||||
return ParsedTranscript{}, err
|
||||
}
|
||||
return ParsedTranscript{
|
||||
Shape: ShapeBareSegmentsArray,
|
||||
Segments: segments,
|
||||
Shape: ShapeBareSegmentsArray,
|
||||
InputSegmentCount: inputCount,
|
||||
Stats: stats,
|
||||
Segments: segments,
|
||||
}, nil
|
||||
default:
|
||||
return ParsedTranscript{}, fmt.Errorf("normalize input must be a top-level object with \"segments\" or a top-level segment array")
|
||||
@@ -121,72 +137,73 @@ func parseObjectShape(raw []byte) (ParsedTranscript, error) {
|
||||
return ParsedTranscript{}, fmt.Errorf("normalize object input must contain a \"segments\" field")
|
||||
}
|
||||
|
||||
segments, err := parseSegmentsArray(segmentsRaw)
|
||||
segments, stats, inputCount, err := parseSegmentsArray(segmentsRaw)
|
||||
if err != nil {
|
||||
return ParsedTranscript{}, err
|
||||
}
|
||||
|
||||
return ParsedTranscript{
|
||||
Shape: ShapeObjectWithSegments,
|
||||
Segments: segments,
|
||||
Shape: ShapeObjectWithSegments,
|
||||
InputSegmentCount: inputCount,
|
||||
Stats: stats,
|
||||
Segments: segments,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func parseSegmentsArray(raw []byte) ([]InputSegment, error) {
|
||||
func parseSegmentsArray(raw []byte) ([]InputSegment, NormalizeStats, int, error) {
|
||||
var segmentValues []json.RawMessage
|
||||
if err := json.Unmarshal(raw, &segmentValues); err != nil {
|
||||
return nil, fmt.Errorf("normalize input \"segments\" must be an array")
|
||||
return nil, NormalizeStats{}, 0, fmt.Errorf("normalize input \"segments\" must be an array")
|
||||
}
|
||||
|
||||
segments := make([]InputSegment, len(segmentValues))
|
||||
for index, segmentRaw := range segmentValues {
|
||||
segment, err := parseSegment(index, segmentRaw)
|
||||
segment, err := decodeSegment(index, segmentRaw)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
return nil, NormalizeStats{}, 0, err
|
||||
}
|
||||
segments[index] = segment
|
||||
}
|
||||
return segments, nil
|
||||
normalized, stats, err := normalizeSegments(segments)
|
||||
return normalized, stats, len(segmentValues), err
|
||||
}
|
||||
|
||||
func parseSegment(index int, raw []byte) (InputSegment, error) {
|
||||
func decodeSegment(index int, raw []byte) (InputSegment, error) {
|
||||
var payload inputSegmentPayload
|
||||
if err := json.Unmarshal(raw, &payload); err != nil {
|
||||
return InputSegment{}, fmt.Errorf("segment %d: invalid segment object: %w", index, err)
|
||||
}
|
||||
|
||||
if payload.Start == nil {
|
||||
return InputSegment{}, fmt.Errorf("segment %d is missing required field \"start\"", index)
|
||||
var start float64
|
||||
if payload.Start != nil {
|
||||
start = *payload.Start
|
||||
}
|
||||
if payload.End == nil {
|
||||
return InputSegment{}, fmt.Errorf("segment %d is missing required field \"end\"", index)
|
||||
}
|
||||
if payload.Speaker == nil {
|
||||
return InputSegment{}, fmt.Errorf("segment %d is missing required field \"speaker\"", index)
|
||||
}
|
||||
if payload.Text == nil {
|
||||
return InputSegment{}, fmt.Errorf("segment %d is missing required field \"text\"", index)
|
||||
var end float64
|
||||
if payload.End != nil {
|
||||
end = *payload.End
|
||||
}
|
||||
|
||||
if *payload.Start < 0 {
|
||||
return InputSegment{}, fmt.Errorf("segment %d has start %v; start must be >= 0", index, *payload.Start)
|
||||
}
|
||||
if *payload.End < *payload.Start {
|
||||
return InputSegment{}, fmt.Errorf("segment %d has end %v before start %v", index, *payload.End, *payload.Start)
|
||||
speaker := ""
|
||||
if payload.Speaker != nil {
|
||||
speaker = strings.TrimSpace(*payload.Speaker)
|
||||
}
|
||||
|
||||
speaker := strings.TrimSpace(*payload.Speaker)
|
||||
if speaker == "" {
|
||||
return InputSegment{}, fmt.Errorf("segment %d has empty \"speaker\"; speaker must be non-empty", index)
|
||||
text := ""
|
||||
if payload.Text != nil {
|
||||
text = *payload.Text
|
||||
}
|
||||
|
||||
return InputSegment{
|
||||
InputIndex: index,
|
||||
OriginalID: payload.ID,
|
||||
Start: *payload.Start,
|
||||
End: *payload.End,
|
||||
StartPresent: payload.Start != nil,
|
||||
EndPresent: payload.End != nil,
|
||||
SpeakerPresent: payload.Speaker != nil,
|
||||
TextPresent: payload.Text != nil,
|
||||
Start: start,
|
||||
End: end,
|
||||
Speaker: speaker,
|
||||
Text: *payload.Text,
|
||||
Text: text,
|
||||
Categories: append([]string(nil), payload.Categories...),
|
||||
Source: payload.Source,
|
||||
SourceSegmentIndex: payload.SourceSegmentIndex,
|
||||
@@ -195,3 +212,96 @@ func parseSegment(index int, raw []byte) (InputSegment, error) {
|
||||
OverlapGroupID: payload.OverlapGroupID,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func normalizeSegments(segments []InputSegment) ([]InputSegment, NormalizeStats, error) {
|
||||
stats := NormalizeStats{}
|
||||
filtered := make([]InputSegment, 0, len(segments))
|
||||
for _, segment := range segments {
|
||||
if strings.TrimSpace(segment.Text) == "" {
|
||||
stats.SegmentsDroppedText++
|
||||
continue
|
||||
}
|
||||
filtered = append(filtered, segment)
|
||||
}
|
||||
|
||||
if len(filtered) == 0 {
|
||||
return filtered, stats, nil
|
||||
}
|
||||
|
||||
hasStart := make([]bool, len(filtered))
|
||||
hasEnd := make([]bool, len(filtered))
|
||||
for index, segment := range filtered {
|
||||
hasStart[index] = segment.StartPresent
|
||||
hasEnd[index] = segment.EndPresent
|
||||
}
|
||||
for index := range filtered {
|
||||
if !hasStart[index] && !hasEnd[index] {
|
||||
midpoint := inferTimestampFromNeighbors(filtered, hasStart, hasEnd, index)
|
||||
filtered[index].Start = midpoint
|
||||
filtered[index].End = midpoint
|
||||
stats.TimingFieldsRepaired += 2
|
||||
hasStart[index] = true
|
||||
hasEnd[index] = true
|
||||
continue
|
||||
}
|
||||
if hasStart[index] && !hasEnd[index] {
|
||||
filtered[index].End = filtered[index].Start
|
||||
stats.TimingFieldsRepaired++
|
||||
hasEnd[index] = true
|
||||
}
|
||||
if !hasStart[index] && hasEnd[index] {
|
||||
filtered[index].Start = filtered[index].End
|
||||
stats.TimingFieldsRepaired++
|
||||
hasStart[index] = true
|
||||
}
|
||||
}
|
||||
|
||||
for index := range filtered {
|
||||
if filtered[index].Start < 0 {
|
||||
return nil, NormalizeStats{}, fmt.Errorf("segment %d has start %v; start must be >= 0", filtered[index].InputIndex, filtered[index].Start)
|
||||
}
|
||||
if filtered[index].End < filtered[index].Start {
|
||||
filtered[index].Start, filtered[index].End = filtered[index].End, filtered[index].Start
|
||||
stats.TimingOrderSwapped++
|
||||
}
|
||||
if strings.TrimSpace(filtered[index].Speaker) == "" {
|
||||
filtered[index].Speaker = "Unknown_Speaker"
|
||||
stats.SpeakerFilled++
|
||||
}
|
||||
}
|
||||
|
||||
return filtered, stats, nil
|
||||
}
|
||||
|
||||
func inferTimestampFromNeighbors(segments []InputSegment, hasStart []bool, hasEnd []bool, index int) float64 {
|
||||
prevEnd, hasPrev := nearestPreviousEnd(segments, hasEnd, index)
|
||||
nextStart, hasNext := nearestNextStart(segments, hasStart, index)
|
||||
if hasPrev && hasNext {
|
||||
return (prevEnd + nextStart) / 2
|
||||
}
|
||||
if hasPrev {
|
||||
return prevEnd
|
||||
}
|
||||
if hasNext {
|
||||
return nextStart
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
func nearestPreviousEnd(segments []InputSegment, hasEnd []bool, index int) (float64, bool) {
|
||||
for i := index - 1; i >= 0; i-- {
|
||||
if hasEnd[i] {
|
||||
return segments[i].End, true
|
||||
}
|
||||
}
|
||||
return 0, false
|
||||
}
|
||||
|
||||
func nearestNextStart(segments []InputSegment, hasStart []bool, index int) (float64, bool) {
|
||||
for i := index + 1; i < len(segments); i++ {
|
||||
if hasStart[i] {
|
||||
return segments[i].Start, true
|
||||
}
|
||||
}
|
||||
return 0, false
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user