Updated the normalize command to correct common errors in WhisperX-generated input transcripts
All checks were successful
ci/woodpecker/tag/release Pipeline was successful

This commit is contained in:
2026-05-16 23:05:42 -05:00
parent 6dbb7ab17e
commit b20438acf0
7 changed files with 357 additions and 66 deletions

View File

@@ -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
}